aboutsummaryrefslogtreecommitdiff
path: root/tests/fixtures/models.py
diff options
context:
space:
mode:
Diffstat (limited to 'tests/fixtures/models.py')
-rw-r--r--tests/fixtures/models.py68
1 files changed, 31 insertions, 37 deletions
diff --git a/tests/fixtures/models.py b/tests/fixtures/models.py
index 671c7b3..157daec 100644
--- a/tests/fixtures/models.py
+++ b/tests/fixtures/models.py
@@ -1,23 +1,22 @@
-from datetime import timezone, datetime
+from collections.abc import Callable, Generator
+from datetime import datetime, timedelta, timezone
+from typing import TYPE_CHECKING
import pytest
+from generalresearch.currency import USDCent
+from generalresearch.pg_helper import PostgresConfig
+from psycopg.errors import ForeignKeyViolation
-from jb.models.event import MTurkEvent
-from generalresearchutils.pg_helper import PostgresConfig
-
-from datetime import datetime, timezone, timedelta
-from typing import Optional, TYPE_CHECKING, Callable, Generator
from jb.managers.amt import AMTManager
-from jb.models.assignment import AssignmentStub, Assignment
-from generalresearchutils.currency import USDCent
-from jb.models.definitions import HitStatus, HitReviewStatus, AssignmentStatus
-from jb.models.hit import HitType, HitQuestion, Hit
+from jb.models.assignment import Assignment, AssignmentStub
+from jb.models.definitions import AssignmentStatus, HitReviewStatus, HitStatus
+from jb.models.event import MTurkEvent
+from jb.models.hit import Hit, HitQuestion, HitType
from tests import generate_amt_id
-from psycopg.errors import ForeignKeyViolation
if TYPE_CHECKING:
- from jb.managers.hit import HitQuestionManager, HitTypeManager, HitManager
from jb.managers.assignment import AssignmentManager
+ from jb.managers.hit import HitManager, HitQuestionManager, HitTypeManager
# --- MTurk Event ---
@@ -76,10 +75,9 @@ def hit_type_record(
yield ht
try:
- with pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute("DELETE FROM mtwerk_hittype WHERE id=%s", (ht.id,))
- conn.commit()
+ with pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute("DELETE FROM mtwerk_hittype WHERE id=%s", (ht.id,))
+ conn.commit()
except ForeignKeyViolation:
pass # DB gets dropped anyway, don't care
@@ -99,10 +97,9 @@ def hit_type_record_with_amt_id(
yield ht
try:
- with pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute("DELETE FROM mtwerk_hittype WHERE id=%s", (ht.id,))
- conn.commit()
+ with pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute("DELETE FROM mtwerk_hittype WHERE id=%s", (ht.id,))
+ conn.commit()
except ForeignKeyViolation:
pass # DB gets dropped anyway, don't care
@@ -166,10 +163,9 @@ def hit_record(
yield hit
try:
- with pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute("DELETE FROM mtwerk_hit WHERE id=%s", (hit.id,))
- conn.commit()
+ with pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute("DELETE FROM mtwerk_hit WHERE id=%s", (hit.id,))
+ conn.commit()
except ForeignKeyViolation:
pass # DB gets dropped anyway, don't care
@@ -228,12 +224,11 @@ def assignment_stub_record(
yield assignment_stub
try:
- with pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(
- "DELETE FROM mtwerk_assignment WHERE id=%s", (assignment_stub.id,)
- )
- conn.commit()
+ with pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(
+ "DELETE FROM mtwerk_assignment WHERE id=%s", (assignment_stub.id,)
+ )
+ conn.commit()
except ForeignKeyViolation:
pass # DB gets dropped anyway, don't care
@@ -255,19 +250,18 @@ def assignment_record(
yield assignment
try:
- with pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute("DELETE FROM mtwerk_assignment WHERE id=%s", (assignment.id,))
- conn.commit()
+ with pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute("DELETE FROM mtwerk_assignment WHERE id=%s", (assignment.id,))
+ conn.commit()
except ForeignKeyViolation:
pass # DB gets dropped anyway, don't care
@pytest.fixture
-def assignment_factory(hit_record: Hit) -> Callable[[Optional[str]], Assignment]:
+def assignment_factory(hit_record: Hit) -> Callable[[str | None], Assignment]:
- def _inner(amt_worker_id: Optional[str] = None) -> Assignment:
+ def _inner(amt_worker_id: str | None = None) -> Assignment:
now = datetime.now(tz=timezone.utc)
amt_assignment_id = generate_amt_id()
amt_worker_id = amt_worker_id or generate_amt_id()
@@ -292,7 +286,7 @@ def assignment_record_factory(
am: "AssignmentManager", assignment_factory: Callable[..., Assignment]
) -> Callable[..., Assignment]:
- def _inner(hit_id: int, amt_worker_id: Optional[str] = None) -> Assignment:
+ def _inner(hit_id: int, amt_worker_id: str | None = None) -> Assignment:
a = assignment_factory(amt_worker_id=amt_worker_id)
a.hit_id = hit_id
am.create_stub(a)