From f094565491e02a916bd32e88977edfd9768d2268 Mon Sep 17 00:00:00 2001 From: Brandur Date: Sat, 6 Jul 2024 14:51:14 -0700 Subject: [PATCH] Fix a problem where tests were leaving leftover jobs in the database Fix a problem where tests were leaving leftover jobs in the database which turns out to be because SQLAlchemy stupidly doesn't start a transaction when `begin()` is called, which was having the effect of the unique job check starting a top-level transaction instead of savepoint, thereby committing insert results outside of the test transactions. This took an insanely long time to track down, and I ended up adding a number of extra measures to help find the problem which I'm also keeping for the time begin: * `RIVER_DEBUG` activates debug logging for SQLAlchemy, showing queries and connection pool activity. * An autorunning fixture checks for extra jobs after each test case to make sure nothing gets leftover. * `SimpleArgs` starts to internalize the name of the test that created it, so upon seeing an extraneous row, it's easy to know exactly where it came from. It becomes a fixture to make this terse. --- README.md | 2 +- .../riversqlalchemy/dbsqlc/river_job.py | 52 ++++- .../riversqlalchemy/dbsqlc/river_job.sql | 4 + tests/client_test.py | 44 ++-- tests/conftest.py | 67 +++++- .../riversqlalchemy/sqlalchemy_driver_test.py | 215 ++++++++++-------- tests/simple_args.py | 11 - 7 files changed, 263 insertions(+), 132 deletions(-) delete mode 100644 tests/simple_args.py diff --git a/README.md b/README.md index adffd84..f464ad4 100644 --- a/README.md +++ b/README.md @@ -97,7 +97,7 @@ insert_res.unique_skipped_as_duplicated ### Custom advisory lock prefix -Unique job insertion takes a Postgres advisory lock to make sure that it's uniqueness check still works even if two conflicting insert operations are occurring in parallel. Postgres advisory locks share a global 64-bit namespace, which is a large enough space that it's unlikely for two advisory locks to ever conflict, but to _guarantee_ that River's advisory locks never interfere with an application's, River can be configured with a 32-bit advisory lock prefix which it will use for all its locks: +Unique job insertion takes a Postgres advisory lock to make sure that its uniqueness check still works even if two conflicting insert operations are occurring in parallel. Postgres advisory locks share a global 64-bit namespace, which is a large enough space that it's unlikely for two advisory locks to ever conflict, but to _guarantee_ that River's advisory locks never interfere with an application's, River can be configured with a 32-bit advisory lock prefix which it will use for all its locks: ```python client = riverqueue.Client(riversqlalchemy.Driver(engine), advisory_lock_prefix: 123456) diff --git a/src/riverqueue/driver/riversqlalchemy/dbsqlc/river_job.py b/src/riverqueue/driver/riversqlalchemy/dbsqlc/river_job.py index d18a5ae..80d098e 100644 --- a/src/riverqueue/driver/riversqlalchemy/dbsqlc/river_job.py +++ b/src/riverqueue/driver/riversqlalchemy/dbsqlc/river_job.py @@ -4,7 +4,7 @@ # source: river_job.sql import dataclasses import datetime -from typing import Any, List, Optional +from typing import Any, AsyncIterator, Iterator, List, Optional import sqlalchemy import sqlalchemy.ext.asyncio @@ -12,6 +12,12 @@ from . import models +JOB_GET_ALL = """-- name: job_get_all \\:many +SELECT id, args, attempt, attempted_at, attempted_by, created_at, errors, finalized_at, kind, max_attempts, metadata, priority, queue, state, scheduled_at, tags +FROM river_job +""" + + JOB_GET_BY_KIND_AND_UNIQUE_PROPERTIES = """-- name: job_get_by_kind_and_unique_properties \\:one SELECT id, args, attempt, attempted_at, attempted_by, created_at, errors, finalized_at, kind, max_attempts, metadata, priority, queue, state, scheduled_at, tags FROM river_job @@ -125,6 +131,28 @@ class Querier: def __init__(self, conn: sqlalchemy.engine.Connection): self._conn = conn + def job_get_all(self) -> Iterator[models.RiverJob]: + result = self._conn.execute(sqlalchemy.text(JOB_GET_ALL)) + for row in result: + yield models.RiverJob( + id=row[0], + args=row[1], + attempt=row[2], + attempted_at=row[3], + attempted_by=row[4], + created_at=row[5], + errors=row[6], + finalized_at=row[7], + kind=row[8], + max_attempts=row[9], + metadata=row[10], + priority=row[11], + queue=row[12], + state=row[13], + scheduled_at=row[14], + tags=row[15], + ) + def job_get_by_kind_and_unique_properties(self, arg: JobGetByKindAndUniquePropertiesParams) -> Optional[models.RiverJob]: row = self._conn.execute(sqlalchemy.text(JOB_GET_BY_KIND_AND_UNIQUE_PROPERTIES), { "p1": arg.kind, @@ -213,6 +241,28 @@ class AsyncQuerier: def __init__(self, conn: sqlalchemy.ext.asyncio.AsyncConnection): self._conn = conn + async def job_get_all(self) -> AsyncIterator[models.RiverJob]: + result = await self._conn.stream(sqlalchemy.text(JOB_GET_ALL)) + async for row in result: + yield models.RiverJob( + id=row[0], + args=row[1], + attempt=row[2], + attempted_at=row[3], + attempted_by=row[4], + created_at=row[5], + errors=row[6], + finalized_at=row[7], + kind=row[8], + max_attempts=row[9], + metadata=row[10], + priority=row[11], + queue=row[12], + state=row[13], + scheduled_at=row[14], + tags=row[15], + ) + async def job_get_by_kind_and_unique_properties(self, arg: JobGetByKindAndUniquePropertiesParams) -> Optional[models.RiverJob]: row = (await self._conn.execute(sqlalchemy.text(JOB_GET_BY_KIND_AND_UNIQUE_PROPERTIES), { "p1": arg.kind, diff --git a/src/riverqueue/driver/riversqlalchemy/dbsqlc/river_job.sql b/src/riverqueue/driver/riversqlalchemy/dbsqlc/river_job.sql index cff1bcd..dd0402d 100644 --- a/src/riverqueue/driver/riversqlalchemy/dbsqlc/river_job.sql +++ b/src/riverqueue/driver/riversqlalchemy/dbsqlc/river_job.sql @@ -35,6 +35,10 @@ CREATE TABLE river_job( CONSTRAINT kind_length CHECK (char_length(kind) > 0 AND char_length(kind) < 128) ); +-- name: JobGetAll :many +SELECT * +FROM river_job; + -- name: JobGetByKindAndUniqueProperties :one SELECT * FROM river_job diff --git a/tests/client_test.py b/tests/client_test.py index 6cd3716..705f980 100644 --- a/tests/client_test.py +++ b/tests/client_test.py @@ -8,8 +8,6 @@ from riverqueue.driver import DriverProtocol, ExecutorProtocol import sqlalchemy -from tests.simple_args import SimpleArgs - @pytest.fixture def mock_driver() -> DriverProtocol: @@ -38,17 +36,17 @@ def client(mock_driver) -> Client: return Client(mock_driver) -def test_insert_with_only_args(client, mock_exec): +def test_insert_with_only_args(client, mock_exec, simple_args): mock_exec.job_get_by_kind_and_unique_properties.return_value = None mock_exec.job_insert.return_value = "job_row" - insert_res = client.insert(SimpleArgs()) + insert_res = client.insert(simple_args) mock_exec.job_insert.assert_called_once() assert insert_res.job == "job_row" -def test_insert_tx(mock_driver, client): +def test_insert_tx(mock_driver, client, simple_args): mock_exec = MagicMock(spec=ExecutorProtocol) mock_exec.job_get_by_kind_and_unique_properties.return_value = None mock_exec.job_insert.return_value = "job_row" @@ -61,17 +59,17 @@ def mock_unwrap_executor(tx: sqlalchemy.Transaction): mock_driver.unwrap_executor.side_effect = mock_unwrap_executor - insert_res = client.insert_tx(mock_tx, SimpleArgs()) + insert_res = client.insert_tx(mock_tx, simple_args) mock_exec.job_insert.assert_called_once() assert insert_res.job == "job_row" -def test_insert_with_insert_opts_from_args(client, mock_exec): +def test_insert_with_insert_opts_from_args(client, mock_exec, simple_args): mock_exec.job_insert.return_value = "job_row" insert_res = client.insert( - SimpleArgs(), + simple_args, insert_opts=InsertOpts( max_attempts=23, priority=2, queue="job_custom_queue", tags=["job_custom"] ), @@ -121,7 +119,7 @@ def to_json() -> str: assert insert_args.tags == ["job_custom"] -def test_insert_with_insert_opts_precedence(client, mock_exec): +def test_insert_with_insert_opts_precedence(client, mock_exec, simple_args): @dataclass class MyArgs: kind = "my_args" @@ -142,7 +140,7 @@ def to_json() -> str: mock_exec.job_insert.return_value = "job_row" insert_res = client.insert( - SimpleArgs(), + simple_args, insert_opts=InsertOpts( max_attempts=17, priority=3, queue="my_queue", tags=["custom"] ), @@ -158,13 +156,13 @@ def to_json() -> str: assert insert_args.tags == ["custom"] -def test_insert_with_unique_opts_by_args(client, mock_exec): +def test_insert_with_unique_opts_by_args(client, mock_exec, simple_args): insert_opts = InsertOpts(unique_opts=UniqueOpts(by_args=True)) mock_exec.job_get_by_kind_and_unique_properties.return_value = None mock_exec.job_insert.return_value = "job_row" - insert_res = client.insert(SimpleArgs(), insert_opts=insert_opts) + insert_res = client.insert(simple_args, insert_opts=insert_opts) mock_exec.job_insert.assert_called_once() assert insert_res.job == "job_row" @@ -175,7 +173,9 @@ def test_insert_with_unique_opts_by_args(client, mock_exec): @patch("datetime.datetime") -def test_insert_with_unique_opts_by_period(mock_datetime, client, mock_exec): +def test_insert_with_unique_opts_by_period( + mock_datetime, client, mock_exec, simple_args +): mock_datetime.now.return_value = datetime(2024, 6, 1, 12, 0, 0, tzinfo=timezone.utc) insert_opts = InsertOpts(unique_opts=UniqueOpts(by_period=900)) @@ -183,7 +183,7 @@ def test_insert_with_unique_opts_by_period(mock_datetime, client, mock_exec): mock_exec.job_get_by_kind_and_unique_properties.return_value = None mock_exec.job_insert.return_value = "job_row" - insert_res = client.insert(SimpleArgs(), insert_opts=insert_opts) + insert_res = client.insert(simple_args, insert_opts=insert_opts) mock_exec.job_insert.assert_called_once() assert insert_res.job == "job_row" @@ -193,13 +193,13 @@ def test_insert_with_unique_opts_by_period(mock_datetime, client, mock_exec): assert call_args.kind == "simple" -def test_insert_with_unique_opts_by_queue(client, mock_exec): +def test_insert_with_unique_opts_by_queue(client, mock_exec, simple_args): insert_opts = InsertOpts(unique_opts=UniqueOpts(by_queue=True)) mock_exec.job_get_by_kind_and_unique_properties.return_value = None mock_exec.job_insert.return_value = "job_row" - insert_res = client.insert(SimpleArgs(), insert_opts=insert_opts) + insert_res = client.insert(simple_args, insert_opts=insert_opts) mock_exec.job_insert.assert_called_once() assert insert_res.job == "job_row" @@ -209,13 +209,13 @@ def test_insert_with_unique_opts_by_queue(client, mock_exec): assert call_args.kind == "simple" -def test_insert_with_unique_opts_by_state(client, mock_exec): +def test_insert_with_unique_opts_by_state(client, mock_exec, simple_args): insert_opts = InsertOpts(unique_opts=UniqueOpts(by_state=["available", "running"])) mock_exec.job_get_by_kind_and_unique_properties.return_value = None mock_exec.job_insert.return_value = "job_row" - insert_res = client.insert(SimpleArgs(), insert_opts=insert_opts) + insert_res = client.insert(simple_args, insert_opts=insert_opts) mock_exec.job_insert.assert_called_once() assert insert_res.job == "job_row" @@ -259,20 +259,20 @@ def to_json() -> None: assert "args should return non-nil from `to_json`" == str(ex.value) -def test_tag_validation(client): +def test_tag_validation(client, simple_args): client.insert( - SimpleArgs(), insert_opts=InsertOpts(tags=["foo", "bar", "baz", "foo-bar-baz"]) + simple_args, insert_opts=InsertOpts(tags=["foo", "bar", "baz", "foo-bar-baz"]) ) with pytest.raises(AssertionError) as ex: - client.insert(SimpleArgs(), insert_opts=InsertOpts(tags=["commas,bad"])) + client.insert(simple_args, insert_opts=InsertOpts(tags=["commas,bad"])) assert ( r"tags should be less than 255 characters in length and match regex \A[\w][\w\-]+[\w]\Z" == str(ex.value) ) with pytest.raises(AssertionError) as ex: - client.insert(SimpleArgs(), insert_opts=InsertOpts(tags=["a" * 256])) + client.insert(simple_args, insert_opts=InsertOpts(tags=["a" * 256])) assert ( r"tags should be less than 255 characters in length and match regex \A[\w][\w\-]+[\w]\Z" == str(ex.value) diff --git a/tests/conftest.py b/tests/conftest.py index ccc855e..2fdffcb 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,13 +1,30 @@ +from dataclasses import dataclass +import json import os +from typing import Iterator import pytest import sqlalchemy import sqlalchemy.ext.asyncio +from riverqueue.driver.riversqlalchemy.dbsqlc import river_job + + +def engine_opts() -> dict: + """ + Use to pass verbose logging options to an SQLAlchemy when `RIVER_DEBUG=true` + in the environment. + """ + + if os.getenv("RIVER_DEBUG") == "true": + return dict(echo=True, echo_pool="debug") + + return dict() + @pytest.fixture(scope="session") def engine() -> sqlalchemy.Engine: - return sqlalchemy.create_engine(test_database_url()) + return sqlalchemy.create_engine(test_database_url(), **engine_opts()) # @pytest_asyncio.fixture(scope="session") @@ -22,10 +39,17 @@ def engine_async() -> sqlalchemy.ext.asyncio.AsyncEngine: # This statement disables pooling which isn't ideal, but I've spent # too many hours trying to figure this out so I'm calling it. poolclass=sqlalchemy.pool.NullPool, + **engine_opts(), ) def test_database_url(is_async: bool = False) -> str: + """ + Produces a test URL based on `TEST_DATABASE_URL` or River's default + convention and modifies it so that it's protocol includes an appropriate + driver to make SQLAlchemy happy. + """ + database_url = os.getenv("TEST_DATABASE_URL", "postgres://localhost/river_test") # sqlalchemy removed support for postgres:// for reasons beyond comprehension @@ -36,3 +60,44 @@ def test_database_url(is_async: bool = False) -> str: database_url = database_url.replace("postgresql://", "postgresql+asyncpg://") return database_url + + +@dataclass +class SimpleArgs: + test_name: str + + kind: str = "simple" + + def to_json(self) -> str: + return json.dumps({"test_name": self.test_name}) + + +@pytest.fixture +def simple_args(request: pytest.FixtureRequest): + """ + Returns an instance of SimpleArgs encapsulating the running test's name. This + can be useful in cases where a test is accidentally leaving leftovers in the + database. + """ + + return SimpleArgs(test_name=request.node.name) + + +@pytest.fixture(autouse=True) +def check_leftover_jobs(engine) -> Iterator[None]: + """ + Autorunning fixture that checks for leftover jobs after each test case. I + previously had a huge amount of trouble tracking down tests that were + inserting rows despite being in a test transaction and ended up adding this + check, along with naming inserted jobs after their test case. If it turns + these measures haven't been needed in a long time, we can probably remove + them. + """ + + yield + + with engine.begin() as conn_tx: + jobs = river_job.Querier(conn_tx).job_get_all() + assert ( + list(jobs) == [] + ), "test case should not have persisted any jobs after run" diff --git a/tests/driver/riversqlalchemy/sqlalchemy_driver_test.py b/tests/driver/riversqlalchemy/sqlalchemy_driver_test.py index 60d324f..50c5e9a 100644 --- a/tests/driver/riversqlalchemy/sqlalchemy_driver_test.py +++ b/tests/driver/riversqlalchemy/sqlalchemy_driver_test.py @@ -20,31 +20,58 @@ from riverqueue.driver import riversqlalchemy from riverqueue.driver.driver_protocol import JobGetByKindAndUniquePropertiesParam -# from tests.conftest import engine_async -from tests.simple_args import SimpleArgs - @pytest.fixture -def driver(engine: sqlalchemy.Engine) -> Iterator[riversqlalchemy.Driver]: +def test_tx(engine: sqlalchemy.Engine) -> Iterator[sqlalchemy.Connection]: with engine.connect() as conn_tx: - conn_tx.execute(sqlalchemy.text("SET search_path TO public")) - yield riversqlalchemy.Driver(conn_tx) + # Force SQLAlchemy to open a transaction. + # + # SQLAlchemy seems to be designed to operate as surprisingly as + # possible. Invoking `begin()` doesn't actually start a transaction. + # Instead, it only does so lazily when a command is first issued. This + # can be a big problem for our internal code, because when it wants to + # start a transaction of its own to do say, a uniqueness check, unless + # another SQL command has already executed it'll accidentally start a + # top-level transaction instead of one in a test transaction that'll be + # rolled back, and cause our tests to commit test jobs. So to work + # around that, we make sure to fire an initial command, thereby forcing + # a transaction to begin. Absolutely terrible design. + conn_tx.execute(sqlalchemy.text("SELECT 1")) + + yield conn_tx + conn_tx.rollback() +@pytest.fixture +def driver(test_tx: sqlalchemy.Connection) -> riversqlalchemy.Driver: + return riversqlalchemy.Driver(test_tx) + + +@pytest.fixture +def client(driver: riversqlalchemy.Driver) -> Client: + return Client(driver) + + @pytest_asyncio.fixture -async def driver_async( +async def test_tx_async( engine_async: sqlalchemy.ext.asyncio.AsyncEngine, -) -> AsyncIterator[riversqlalchemy.AsyncDriver]: +) -> AsyncIterator[sqlalchemy.ext.asyncio.AsyncConnection]: async with engine_async.connect() as conn_tx: - await conn_tx.execute(sqlalchemy.text("SET search_path TO public")) - yield riversqlalchemy.AsyncDriver(conn_tx) + # Force SQLAlchemy to open a transaction. + # + # See explanatory comment in `test_tx()` above. + await conn_tx.execute(sqlalchemy.text("SELECT 1")) + + yield conn_tx await conn_tx.rollback() @pytest.fixture -def client(driver: riversqlalchemy.Driver) -> Client: - return Client(driver) +def driver_async( + test_tx_async: sqlalchemy.ext.asyncio.AsyncConnection, +) -> riversqlalchemy.AsyncDriver: + return riversqlalchemy.AsyncDriver(test_tx_async) @pytest_asyncio.fixture @@ -54,8 +81,8 @@ async def client_async( return AsyncClient(driver_async) -def test_insert_job_from_row(client, driver): - insert_res = client.insert(SimpleArgs()) +def test_insert_job_from_row(client, simple_args): + insert_res = client.insert(simple_args) job = insert_res.job assert job assert isinstance(job.args, dict) @@ -73,165 +100,163 @@ def test_insert_job_from_row(client, driver): assert job.tags == [] -def test_insert_with_only_args_sync(client, driver): - insert_res = client.insert(SimpleArgs()) +def test_insert_with_only_args_sync(client, simple_args): + insert_res = client.insert(simple_args) assert insert_res.job @pytest.mark.asyncio -async def test_insert_with_only_args_async(client_async): - insert_res = await client_async.insert(SimpleArgs()) +async def test_insert_with_only_args_async(client_async, simple_args): + insert_res = await client_async.insert(simple_args) assert insert_res.job -def test_insert_tx_sync(client, driver, engine): - with engine.begin() as conn_tx: - args = SimpleArgs() - insert_res = client.insert_tx(conn_tx, args) - assert insert_res.job +def test_insert_tx_sync(client, driver, engine, simple_args, test_tx): + insert_res = client.insert_tx(test_tx, simple_args) + assert insert_res.job - job = driver.unwrap_executor(conn_tx).job_get_by_kind_and_unique_properties( - JobGetByKindAndUniquePropertiesParam(kind=args.kind) + job = driver.unwrap_executor(test_tx).job_get_by_kind_and_unique_properties( + JobGetByKindAndUniquePropertiesParam(kind=simple_args.kind) + ) + assert job == insert_res.job + + with engine.begin() as conn_tx2: + job = driver.unwrap_executor(conn_tx2).job_get_by_kind_and_unique_properties( + JobGetByKindAndUniquePropertiesParam(kind=simple_args.kind) ) - assert job == insert_res.job + assert job is None - with engine.begin() as conn_tx2: - job = driver.unwrap_executor( - conn_tx2 - ).job_get_by_kind_and_unique_properties( - JobGetByKindAndUniquePropertiesParam(kind=args.kind) - ) - assert job is None - - conn_tx.rollback() + conn_tx2.rollback() @pytest.mark.asyncio -async def test_insert_tx_async(client_async, driver_async, engine_async): - async with engine_async.begin() as conn_tx: - args = SimpleArgs() - insert_res = await client_async.insert_tx(conn_tx, args) - assert insert_res.job +async def test_insert_tx_async( + client_async, driver_async, engine_async, simple_args, test_tx_async +): + insert_res = await client_async.insert_tx(test_tx_async, simple_args) + assert insert_res.job + job = await driver_async.unwrap_executor( + test_tx_async + ).job_get_by_kind_and_unique_properties( + JobGetByKindAndUniquePropertiesParam(kind=simple_args.kind) + ) + assert job == insert_res.job + + async with engine_async.begin() as conn_tx2: job = await driver_async.unwrap_executor( - conn_tx + conn_tx2 ).job_get_by_kind_and_unique_properties( - JobGetByKindAndUniquePropertiesParam(kind=args.kind) + JobGetByKindAndUniquePropertiesParam(kind=simple_args.kind) ) - assert job == insert_res.job - - async with engine_async.begin() as conn_tx2: - job = await driver_async.unwrap_executor( - conn_tx2 - ).job_get_by_kind_and_unique_properties( - JobGetByKindAndUniquePropertiesParam(kind=args.kind) - ) - assert job is None + assert job is None - await conn_tx.rollback() + await conn_tx2.rollback() -def test_insert_with_opts_sync(client): +def test_insert_with_opts_sync(client, simple_args): insert_opts = InsertOpts(queue="high_priority", unique_opts=None) - insert_res = client.insert(SimpleArgs(), insert_opts=insert_opts) + insert_res = client.insert(simple_args, insert_opts=insert_opts) assert insert_res.job @pytest.mark.asyncio -async def test_insert_with_opts_async(client_async): +async def test_insert_with_opts_async(client_async, simple_args): insert_opts = InsertOpts(queue="high_priority", unique_opts=None) - insert_res = await client_async.insert(SimpleArgs(), insert_opts=insert_opts) + insert_res = await client_async.insert(simple_args, insert_opts=insert_opts) assert insert_res.job -def test_insert_with_unique_opts_by_args_sync(client): +def test_insert_with_unique_opts_by_args_sync(client, simple_args): insert_opts = InsertOpts(unique_opts=UniqueOpts(by_args=True)) - insert_res = client.insert(SimpleArgs(), insert_opts=insert_opts) + insert_res = client.insert(simple_args, insert_opts=insert_opts) assert insert_res.job - insert_res2 = client.insert(SimpleArgs(), insert_opts=insert_opts) + insert_res2 = client.insert(simple_args, insert_opts=insert_opts) assert insert_res.job == insert_res2.job @pytest.mark.asyncio -async def test_insert_with_unique_opts_by_args_sync_async(client_async): +async def test_insert_with_unique_opts_by_args_async(client_async, simple_args): insert_opts = InsertOpts(unique_opts=UniqueOpts(by_args=True)) - insert_res = await client_async.insert(SimpleArgs(), insert_opts=insert_opts) + insert_res = await client_async.insert(simple_args, insert_opts=insert_opts) assert insert_res.job - insert_res2 = await client_async.insert(SimpleArgs(), insert_opts=insert_opts) + insert_res2 = await client_async.insert(simple_args, insert_opts=insert_opts) assert insert_res.job == insert_res2.job @patch("datetime.datetime") -def test_insert_with_unique_opts_by_period_sync(mock_datetime, client): +def test_insert_with_unique_opts_by_period_sync(mock_datetime, client, simple_args): mock_datetime.now.return_value = datetime(2024, 6, 1, 12, 0, 0, tzinfo=timezone.utc) insert_opts = InsertOpts(unique_opts=UniqueOpts(by_period=900)) - insert_res = client.insert(SimpleArgs(), insert_opts=insert_opts) - insert_res2 = client.insert(SimpleArgs(), insert_opts=insert_opts) + insert_res = client.insert(simple_args, insert_opts=insert_opts) + insert_res2 = client.insert(simple_args, insert_opts=insert_opts) assert insert_res.job == insert_res2.job @patch("datetime.datetime") @pytest.mark.asyncio -async def test_insert_with_unique_opts_by_period_async(mock_datetime, client_async): +async def test_insert_with_unique_opts_by_period_async( + mock_datetime, client_async, simple_args +): mock_datetime.now.return_value = datetime(2024, 6, 1, 12, 0, 0, tzinfo=timezone.utc) insert_opts = InsertOpts(unique_opts=UniqueOpts(by_period=900)) - insert_res = await client_async.insert(SimpleArgs(), insert_opts=insert_opts) - insert_res2 = await client_async.insert(SimpleArgs(), insert_opts=insert_opts) + insert_res = await client_async.insert(simple_args, insert_opts=insert_opts) + insert_res2 = await client_async.insert(simple_args, insert_opts=insert_opts) assert insert_res.job == insert_res2.job -def test_insert_with_unique_opts_by_queue_sync(client): +def test_insert_with_unique_opts_by_queue_sync(client, simple_args): insert_opts = InsertOpts(unique_opts=UniqueOpts(by_queue=True)) - insert_res = client.insert(SimpleArgs(), insert_opts=insert_opts) + insert_res = client.insert(simple_args, insert_opts=insert_opts) assert insert_res.job - insert_res2 = client.insert(SimpleArgs(), insert_opts=insert_opts) + insert_res2 = client.insert(simple_args, insert_opts=insert_opts) assert insert_res.job == insert_res2.job @pytest.mark.asyncio -async def test_insert_with_unique_opts_by_queue_async(client_async): +async def test_insert_with_unique_opts_by_queue_async(client_async, simple_args): insert_opts = InsertOpts(unique_opts=UniqueOpts(by_queue=True)) - insert_res = await client_async.insert(SimpleArgs(), insert_opts=insert_opts) + insert_res = await client_async.insert(simple_args, insert_opts=insert_opts) assert insert_res.job - insert_res2 = await client_async.insert(SimpleArgs(), insert_opts=insert_opts) + insert_res2 = await client_async.insert(simple_args, insert_opts=insert_opts) assert insert_res.job == insert_res2.job -def test_insert_with_unique_opts_by_state_sync(client): +def test_insert_with_unique_opts_by_state_sync(client, simple_args): insert_opts = InsertOpts(unique_opts=UniqueOpts(by_state=["available", "running"])) - insert_res = client.insert(SimpleArgs(), insert_opts=insert_opts) + insert_res = client.insert(simple_args, insert_opts=insert_opts) assert insert_res.job - insert_res2 = client.insert(SimpleArgs(), insert_opts=insert_opts) + insert_res2 = client.insert(simple_args, insert_opts=insert_opts) assert insert_res.job == insert_res2.job @pytest.mark.asyncio -async def test_insert_with_unique_opts_by_state_async(client_async): +async def test_insert_with_unique_opts_by_state_async(client_async, simple_args): insert_opts = InsertOpts(unique_opts=UniqueOpts(by_state=["available", "running"])) - insert_res = await client_async.insert(SimpleArgs(), insert_opts=insert_opts) + insert_res = await client_async.insert(simple_args, insert_opts=insert_opts) assert insert_res.job - insert_res2 = await client_async.insert(SimpleArgs(), insert_opts=insert_opts) + insert_res2 = await client_async.insert(simple_args, insert_opts=insert_opts) assert insert_res.job == insert_res2.job -def test_insert_many_with_only_args_sync(client, driver): - num_inserted = client.insert_many([SimpleArgs()]) +def test_insert_many_with_only_args_sync(client, simple_args): + num_inserted = client.insert_many([simple_args]) assert num_inserted == 1 @pytest.mark.asyncio -async def test_insert_many_with_only_args_async(client_async): - num_inserted = await client_async.insert_many([SimpleArgs()]) +async def test_insert_many_with_only_args_async(client_async, simple_args): + num_inserted = await client_async.insert_many([simple_args]) assert num_inserted == 1 -def test_insert_many_with_insert_opts_sync(client, driver): +def test_insert_many_with_insert_opts_sync(client, simple_args): num_inserted = client.insert_many( [ InsertManyParams( - args=SimpleArgs(), + args=simple_args, insert_opts=InsertOpts(queue="high_priority", unique_opts=None), ) ] @@ -240,11 +265,11 @@ def test_insert_many_with_insert_opts_sync(client, driver): @pytest.mark.asyncio -async def test_insert_many_with_insert_opts_async(client_async): +async def test_insert_many_with_insert_opts_async(client_async, simple_args): num_inserted = await client_async.insert_many( [ InsertManyParams( - args=SimpleArgs(), + args=simple_args, insert_opts=InsertOpts(queue="high_priority", unique_opts=None), ) ] @@ -252,14 +277,12 @@ async def test_insert_many_with_insert_opts_async(client_async): assert num_inserted == 1 -def test_insert_many_tx_sync(client, engine): - with engine.begin() as conn_tx: - num_inserted = client.insert_many_tx(conn_tx, [SimpleArgs()]) - assert num_inserted == 1 +def test_insert_many_tx_sync(client, simple_args, test_tx): + num_inserted = client.insert_many_tx(test_tx, [simple_args]) + assert num_inserted == 1 @pytest.mark.asyncio -async def test_insert_many_tx_async(client_async, engine_async): - async with engine_async.begin() as conn_tx: - num_inserted = await client_async.insert_many_tx(conn_tx, [SimpleArgs()]) - assert num_inserted == 1 +async def test_insert_many_tx_async(client_async, simple_args, test_tx_async): + num_inserted = await client_async.insert_many_tx(test_tx_async, [simple_args]) + assert num_inserted == 1 diff --git a/tests/simple_args.py b/tests/simple_args.py deleted file mode 100644 index 4e89ed4..0000000 --- a/tests/simple_args.py +++ /dev/null @@ -1,11 +0,0 @@ -import json -from dataclasses import dataclass - - -@dataclass -class SimpleArgs: - kind: str = "simple" - - @staticmethod - def to_json() -> str: - return json.dumps({"job_num": 1})