aboutsummaryrefslogtreecommitdiff
path: root/jb
diff options
context:
space:
mode:
Diffstat (limited to 'jb')
-rw-r--r--jb/api/auth.py3
-rw-r--r--jb/api/magic_token.py6
-rw-r--r--jb/flow/maintenance.py4
-rw-r--r--jb/flow/monitoring.py23
-rw-r--r--jb/flow/setup_tasks.py5
-rw-r--r--jb/flow/tasks.py6
-rw-r--r--jb/main.py10
-rw-r--r--jb/managers/__init__.py2
-rw-r--r--jb/managers/amt.py28
-rw-r--r--jb/managers/assignment.py85
-rw-r--r--jb/managers/bonus.py22
-rw-r--r--jb/managers/email_manager.py5
-rw-r--r--jb/managers/gr_api.py6
-rw-r--r--jb/managers/hit.py115
-rw-r--r--jb/managers/thl.py19
-rw-r--r--jb/managers/worker.py6
-rw-r--r--jb/models/__init__.py5
-rw-r--r--jb/models/assignment.py24
-rw-r--r--jb/models/bonus.py12
-rw-r--r--jb/models/custom_types.py7
-rw-r--r--jb/models/errors.py2
-rw-r--r--jb/models/event.py6
-rw-r--r--jb/models/hit.py42
-rw-r--r--jb/settings.py19
24 files changed, 183 insertions, 279 deletions
diff --git a/jb/api/auth.py b/jb/api/auth.py
index a92515d..411b8f1 100644
--- a/jb/api/auth.py
+++ b/jb/api/auth.py
@@ -4,7 +4,7 @@ from typing import Annotated
from uuid import uuid4
import jwt
-from fastapi import Depends, HTTPException, Request, Response, status
+from fastapi import Depends, HTTPException, Request, status
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
from jb.config import settings
@@ -87,4 +87,3 @@ def create_session(product_user_id: str) -> str:
settings.session_jwt_secret.get_secret_value(),
algorithm="HS256",
)
-
diff --git a/jb/api/magic_token.py b/jb/api/magic_token.py
index 562ba81..b7f7575 100644
--- a/jb/api/magic_token.py
+++ b/jb/api/magic_token.py
@@ -4,7 +4,7 @@ import secrets
from fastapi import HTTPException, status
from jb.decorators import REDIS
-from jb.models.auth import AmtAccountLink, User
+from jb.models.auth import AmtAccountLink
MAGIC_TOKEN_PREFIX = "auth:magic:"
AMT_ACCOUNT_LINK_TOKEN_PREFIX = "auth:amt-account-link:"
@@ -60,9 +60,7 @@ def create_amt_account_link_token(email: str, amt_worker_id: str) -> str:
def consume_amt_account_link_token(token: str) -> AmtAccountLink:
"""Atomically consume and validate an AMT account-link token."""
- raw_data = REDIS.getdel(
- redis_token_key(token, AMT_ACCOUNT_LINK_TOKEN_PREFIX)
- )
+ raw_data = REDIS.getdel(redis_token_key(token, AMT_ACCOUNT_LINK_TOKEN_PREFIX))
if raw_data is None:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
diff --git a/jb/flow/maintenance.py b/jb/flow/maintenance.py
index f8ca971..4a2fe81 100644
--- a/jb/flow/maintenance.py
+++ b/jb/flow/maintenance.py
@@ -1,5 +1,3 @@
-from typing import Optional
-
from jb.decorators import HM
from jb.flow.monitoring import emit_hit_event
from jb.managers.amt import AMTManager
@@ -10,7 +8,7 @@ def check_hit_status(
amtm: AMTManager,
amt_hit_id: str,
amt_hit_type_id: str,
- reason: Optional[str] = None,
+ reason: str | None = None,
) -> HitStatus:
"""
(this used to be called "process_hit")
diff --git a/jb/flow/monitoring.py b/jb/flow/monitoring.py
index 334eab3..22e9e93 100644
--- a/jb/flow/monitoring.py
+++ b/jb/flow/monitoring.py
@@ -1,12 +1,11 @@
import socket
-from typing import Optional
+from generalresearch.currency import USDCent
from mypy_boto3_mturk.literals import EventTypeType
from jb.config import settings
from jb.decorators import influx_client
-from generalresearch.currency import USDCent
-from jb.models.definitions import HitStatus, AssignmentStatus
+from jb.models.definitions import AssignmentStatus, HitStatus
def write_hit_gauge(status: HitStatus, amt_hit_type_id: str, cnt: int) -> None:
@@ -26,8 +25,6 @@ def write_hit_gauge(status: HitStatus, amt_hit_type_id: str, cnt: int) -> None:
if influx_client:
influx_client.write_points(points=[point])
- return None
-
def write_assignment_gauge(
status: AssignmentStatus, amt_hit_type_id: str, cnt: int
@@ -47,11 +44,9 @@ def write_assignment_gauge(
if influx_client:
influx_client.write_points(points=[point])
- return None
-
def emit_hit_event(
- status: HitStatus, amt_hit_type_id: str, reason: Optional[str] = None
+ status: HitStatus, amt_hit_type_id: str, reason: str | None = None
) -> None:
"""
e.g. a HIT was created, Reviewable, etc. We don't have a "created"
@@ -77,11 +72,9 @@ def emit_hit_event(
if influx_client:
influx_client.write_points([point])
- return None
-
def emit_assignment_event(
- status: AssignmentStatus, amt_hit_type_id: str, reason: Optional[str] = None
+ status: AssignmentStatus, amt_hit_type_id: str, reason: str | None = None
) -> None:
"""
e.g. an Assignment was accepted/approved/reject
@@ -106,8 +99,6 @@ def emit_assignment_event(
if influx_client:
influx_client.write_points([point])
- return None
-
def emit_mturk_notification_event(
event_type: EventTypeType, amt_hit_type_id: str
@@ -133,8 +124,6 @@ def emit_mturk_notification_event(
if influx_client:
influx_client.write_points([point])
- return None
-
def emit_error_event(event_type: str, amt_hit_type_id: str) -> None:
"""
@@ -157,8 +146,6 @@ def emit_error_event(event_type: str, amt_hit_type_id: str) -> None:
if influx_client:
influx_client.write_points([point])
- return None
-
def emit_bonus_event(amount: USDCent, amt_hit_type_id: str) -> None:
"""
@@ -179,5 +166,3 @@ def emit_bonus_event(amount: USDCent, amt_hit_type_id: str) -> None:
if influx_client:
influx_client.write_points([point])
-
- return None
diff --git a/jb/flow/setup_tasks.py b/jb/flow/setup_tasks.py
index 4664374..f5cda48 100644
--- a/jb/flow/setup_tasks.py
+++ b/jb/flow/setup_tasks.py
@@ -1,6 +1,5 @@
-from jb.config import TOPIC_ARN, SUBSCRIPTION
-from jb.decorators import SNS_CLIENT, AMT_CLIENT
-from jb.config import settings
+from jb.config import SUBSCRIPTION, TOPIC_ARN, settings
+from jb.decorators import AMT_CLIENT, SNS_CLIENT
def initial_setup():
diff --git a/jb/flow/tasks.py b/jb/flow/tasks.py
index 6c5d6bb..24e96d4 100644
--- a/jb/flow/tasks.py
+++ b/jb/flow/tasks.py
@@ -4,11 +4,11 @@ from typing import TypedDict, cast
from generalresearch.config import is_debug
-from jb.decorators import AMTM, HTM, HM, HQM, pg_config
+from jb.decorators import AMTM, HM, HQM, HTM, pg_config
from jb.flow.maintenance import check_hit_status
-from jb.flow.monitoring import write_hit_gauge, emit_hit_event
+from jb.flow.monitoring import emit_hit_event, write_hit_gauge
from jb.models.definitions import HitStatus
-from jb.models.hit import HitType, HitQuestion, Hit
+from jb.models.hit import Hit, HitQuestion, HitType
logging.basicConfig()
logger = logging.getLogger()
diff --git a/jb/main.py b/jb/main.py
index 70c98e9..9f4f000 100644
--- a/jb/main.py
+++ b/jb/main.py
@@ -1,15 +1,15 @@
from multiprocessing import Process
-from typing import Any, Dict
+from typing import Any
from fastapi import FastAPI
from fastapi.responses import HTMLResponse
from starlette.middleware.cors import CORSMiddleware
from starlette.middleware.trustedhost import TrustedHostMiddleware
-from jb.views.common import common_router
-from jb.views.auth import auth_router
-from jb.settings import BASE_HTML
from jb.config import settings
+from jb.settings import BASE_HTML
+from jb.views.auth import auth_router
+from jb.views.common import common_router
app = FastAPI(
servers=[
@@ -38,7 +38,7 @@ app.include_router(router=auth_router)
@app.get("/robots.txt")
@app.get("/sitemap.xml")
@app.get("/favicon.ico")
-def return_nothing() -> Dict[str, Any]:
+def return_nothing() -> dict[str, Any]:
return {}
diff --git a/jb/managers/__init__.py b/jb/managers/__init__.py
index ec64f9f..92ba8bd 100644
--- a/jb/managers/__init__.py
+++ b/jb/managers/__init__.py
@@ -1,5 +1,5 @@
+from collections.abc import Collection
from enum import IntEnum
-from typing import Collection
from generalresearch.pg_helper import PostgresConfig
diff --git a/jb/managers/amt.py b/jb/managers/amt.py
index 2cb0cbd..e2c7e90 100644
--- a/jb/managers/amt.py
+++ b/jb/managers/amt.py
@@ -1,6 +1,6 @@
import logging
-from datetime import timezone, datetime
-from typing import Tuple, Optional, List, Dict, Any
+from datetime import datetime, timezone
+from typing import Any
import botocore.exceptions
from generalresearch.currency import USDCent
@@ -9,8 +9,8 @@ from mypy_boto3_mturk.type_defs import (
AssignmentTypeDef,
BonusPaymentTypeDef,
CreateHITTypeResponseTypeDef,
- GetHITResponseTypeDef,
CreateHITWithHITTypeResponseTypeDef,
+ GetHITResponseTypeDef,
)
from pydantic import ValidationError
@@ -19,7 +19,7 @@ from jb.models import AMTAccount
from jb.models.assignment import Assignment
from jb.models.bonus import Bonus
from jb.models.definitions import HitStatus
-from jb.models.hit import HitType, HitQuestion, Hit
+from jb.models.hit import Hit, HitQuestion, HitType
REJECT_MESSAGE_UNKNOWN_ASSIGNMENT = "Unknown assignment"
REJECT_MESSAGE_NO_WORK = "Assignment was submitted with no attempted work."
@@ -56,7 +56,7 @@ class AMTManager:
}
)
- def get_hit_if_exists(self, amt_hit_id: str) -> Tuple[Optional[Hit], Optional[str]]:
+ def get_hit_if_exists(self, amt_hit_id: str) -> tuple[Hit | None, str | None]:
try:
res: GetHITResponseTypeDef = self.amt_client.get_hit(HITId=amt_hit_id)
@@ -146,7 +146,7 @@ class AMTManager:
assert assignment.id is None
return assignment
- def get_assignment_if_exists(self, amt_assignment_id: str) -> Optional[Assignment]:
+ def get_assignment_if_exists(self, amt_assignment_id: str) -> Assignment | None:
expected_err_msg = f"Assignment {amt_assignment_id} does not exist"
try:
@@ -161,7 +161,7 @@ class AMTManager:
def reject_assignment_if_possible(
self, amt_assignment_id: str, msg: str = REJECT_MESSAGE_UNKNOWN_ASSIGNMENT
- ) -> Optional[Dict[str, Any]]:
+ ) -> dict[str, Any] | None:
# Unclear to me when this would fail
try:
@@ -178,7 +178,7 @@ class AMTManager:
amt_assignment_id: str,
msg: str = APPROVAL_MESSAGE,
override_rejection: bool = False,
- ) -> Optional[Dict[str, Any]]:
+ ) -> dict[str, Any] | None:
# Unclear to me when this would fail
try:
@@ -207,8 +207,6 @@ class AMTManager:
# elif "This HIT is currently in the state 'Reviewing'" in error_msg:
# logging.warning(error_msg)
- return None
-
def send_bonus(
self,
amt_worker_id: str,
@@ -216,7 +214,7 @@ class AMTManager:
amt_assignment_id: str,
reason: str,
unique_request_token: str,
- ) -> Optional[Dict[str, Any]]:
+ ) -> dict[str, Any] | None:
try:
return self.amt_client.send_bonus(
WorkerId=amt_worker_id,
@@ -230,11 +228,9 @@ class AMTManager:
logging.warning(f"{amt_worker_id=} {amt_assignment_id=}, {e}")
return None
- def get_bonus(
- self, amt_assignment_id: str, payout_event_id: str
- ) -> Optional[Bonus]:
+ def get_bonus(self, amt_assignment_id: str, payout_event_id: str) -> Bonus | None:
- res: List[BonusPaymentTypeDef] = self.amt_client.list_bonus_payments(
+ res: list[BonusPaymentTypeDef] = self.amt_client.list_bonus_payments(
AssignmentId=amt_assignment_id
)["BonusPayments"]
@@ -268,5 +264,3 @@ class AMTManager:
self.amt_client.update_expiration_for_hit(
HITId=hit["HITId"], ExpireAt=now
)
-
- return None
diff --git a/jb/managers/assignment.py b/jb/managers/assignment.py
index f6aa2ce..089adb1 100644
--- a/jb/managers/assignment.py
+++ b/jb/managers/assignment.py
@@ -1,11 +1,10 @@
from datetime import datetime, timezone
-from typing import Optional
from psycopg import sql
from pydantic import NonNegativeInt, PositiveInt
from jb.managers import PostgresManager
-from jb.models.assignment import AssignmentStub, Assignment
+from jb.models.assignment import Assignment, AssignmentStub
from jb.models.definitions import AssignmentStatus
@@ -14,8 +13,7 @@ class AssignmentManager(PostgresManager):
def create_stub(self, stub: AssignmentStub) -> None:
assert stub.id is None
data = stub.to_postgres()
- query = sql.SQL(
- """
+ query = sql.SQL("""
INSERT INTO mtwerk_assignment
(amt_assignment_id, amt_worker_id, status,
created_at, modified_at, hit_id)
@@ -23,16 +21,13 @@ class AssignmentManager(PostgresManager):
(%(amt_assignment_id)s, %(amt_worker_id)s, %(status)s,
%(created_at)s, %(modified_at)s, %(hit_id)s)
RETURNING id;
- """
- )
+ """)
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- pk = c.fetchone()["id"] # type: ignore
- conn.commit()
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ pk = c.fetchone()["id"] # type: ignore
+ conn.commit()
stub.id = pk
- return None
def create(self, assignment: Assignment) -> None:
# Typically this is NOT used (we'd create the stub when HIT is
@@ -42,8 +37,7 @@ class AssignmentManager(PostgresManager):
assert assignment.id is None
data = assignment.to_postgres()
- query = sql.SQL(
- """
+ query = sql.SQL("""
INSERT INTO mtwerk_assignment
(amt_assignment_id, amt_worker_id, status,
created_at, modified_at, hit_id,
@@ -57,16 +51,13 @@ class AssignmentManager(PostgresManager):
%(approval_time)s, %(rejection_time)s, %(requester_feedback)s,
%(tsid)s)
RETURNING id;
- """
- )
+ """)
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- pk = c.fetchone()["id"] # type: ignore
- conn.commit()
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ pk = c.fetchone()["id"] # type: ignore
+ conn.commit()
assignment.id = pk
- return None
def get_stub(self, amt_assignment_id: str) -> AssignmentStub:
res = self.pg_config.execute_sql_query(
@@ -85,7 +76,7 @@ class AssignmentManager(PostgresManager):
assert len(res) == 1
return AssignmentStub.model_validate(res[0])
- def get_stub_if_exists(self, amt_assignment_id: str) -> Optional[AssignmentStub]:
+ def get_stub_if_exists(self, amt_assignment_id: str) -> AssignmentStub | None:
try:
return self.get_stub(amt_assignment_id=amt_assignment_id)
except AssertionError:
@@ -121,8 +112,7 @@ class AssignmentManager(PostgresManager):
"amt_assignment_id": assignment.amt_assignment_id,
"modified_at": now,
}
- query = sql.SQL(
- """
+ query = sql.SQL("""
UPDATE mtwerk_assignment
SET submit_time = %(submit_time)s,
auto_approval_time = %(auto_approval_time)s,
@@ -130,15 +120,12 @@ class AssignmentManager(PostgresManager):
tsid = %(tsid)s,
modified_at = %(modified_at)s
WHERE amt_assignment_id = %(amt_assignment_id)s
- """
- )
+ """)
# We force this to fail if the assignment doesn't already exist in the db
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- assert c.rowcount == 1, f"Expected 1 row, got {c.rowcount}"
- conn.commit()
- return None
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ assert c.rowcount == 1, f"Expected 1 row, got {c.rowcount}"
+ conn.commit()
def reject(self, assignment: Assignment) -> None:
assert assignment.status == AssignmentStatus.Rejected
@@ -156,8 +143,7 @@ class AssignmentManager(PostgresManager):
"accept_time": assignment.accept_time,
"modified_at": now,
}
- query = sql.SQL(
- """
+ query = sql.SQL("""
UPDATE mtwerk_assignment
SET submit_time = %(submit_time)s,
rejection_time = %(rejection_time)s,
@@ -167,15 +153,12 @@ class AssignmentManager(PostgresManager):
accept_time = %(accept_time)s,
modified_at = %(modified_at)s
WHERE amt_assignment_id = %(amt_assignment_id)s
- """
- )
+ """)
# We force this to fail if the assignment doesn't already exist in the db
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- assert c.rowcount == 1, f"Expected 1 row, got {c.rowcount}"
- conn.commit()
- return None
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ assert c.rowcount == 1, f"Expected 1 row, got {c.rowcount}"
+ conn.commit()
def approve(self, assignment: Assignment) -> None:
assert assignment.status == AssignmentStatus.Approved
@@ -194,8 +177,7 @@ class AssignmentManager(PostgresManager):
"accept_time": assignment.accept_time,
"modified_at": now,
}
- query = sql.SQL(
- """
+ query = sql.SQL("""
UPDATE mtwerk_assignment
SET submit_time = %(submit_time)s,
approval_time = %(approval_time)s,
@@ -205,15 +187,12 @@ class AssignmentManager(PostgresManager):
accept_time = %(accept_time)s,
modified_at = %(modified_at)s
WHERE amt_assignment_id = %(amt_assignment_id)s
- """
- )
+ """)
# We force this to fail if the assignment doesn't already exist in the db
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- assert c.rowcount == 1, f"Expected 1 row, got {c.rowcount}"
- conn.commit()
- return None
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ assert c.rowcount == 1, f"Expected 1 row, got {c.rowcount}"
+ conn.commit()
def missing_tsid_count(
self, amt_worker_id: str, lookback_hrs: PositiveInt = 24
diff --git a/jb/managers/bonus.py b/jb/managers/bonus.py
index 89b81f0..15d0e5b 100644
--- a/jb/managers/bonus.py
+++ b/jb/managers/bonus.py
@@ -1,4 +1,4 @@
-from typing import List, Any
+from typing import Any
from psycopg import sql
@@ -11,8 +11,7 @@ class BonusManager(PostgresManager):
def create(self, bonus: Bonus) -> None:
assert bonus.id is None
data = bonus.to_postgres()
- query = sql.SQL(
- """
+ query = sql.SQL("""
INSERT INTO mtwerk_bonus
(payout_event_id, amt_worker_id, amount, grant_time, assignment_id, reason)
VALUES (
@@ -29,20 +28,17 @@ class BonusManager(PostgresManager):
%(reason)s
)
RETURNING id, assignment_id;
- """
- )
+ """)
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- res = c.fetchone()
- conn.commit()
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ res = c.fetchone()
+ conn.commit()
bonus.id = res["id"] # type: ignore
bonus.assignment_id = res["assignment_id"] # type: ignore
- return None
- def filter(self, amt_assignment_id: str) -> List[Bonus]:
- res: List[Any] = self.pg_config.execute_sql_query(
+ def filter(self, amt_assignment_id: str) -> list[Bonus]:
+ res: list[Any] = self.pg_config.execute_sql_query(
"""
SELECT mb.*, ma.amt_assignment_id
FROM mtwerk_bonus mb
diff --git a/jb/managers/email_manager.py b/jb/managers/email_manager.py
index dcbe167..e740e20 100644
--- a/jb/managers/email_manager.py
+++ b/jb/managers/email_manager.py
@@ -21,7 +21,7 @@ def get_or_create_contact(email: str, amt_worker_id: str | None = None):
return contact_id
-def send_login_email_from_url(mautic_url, magic_link) -> None:
+def send_login_email_from_url(mautic_url: str, magic_link: str) -> None:
email_tokens = {
"magic_link": magic_link,
}
@@ -48,7 +48,8 @@ def send_login_email(email: str, magic_token: str):
def send_amt_link_email(email: str, magic_token: str):
- # don't actually associate the email with the worker ID until they click the link
+ # Don't actually associate the email with the worker ID
+ # until they click the link
contact_id = get_or_create_contact(email=email)
mautic_url = (
f"{MAUTIC_BASE_URL}/api/emails/{EMAIL_TEMPLATE_ID}/contact/{contact_id}/send"
diff --git a/jb/managers/gr_api.py b/jb/managers/gr_api.py
index 20dc051..494ceea 100644
--- a/jb/managers/gr_api.py
+++ b/jb/managers/gr_api.py
@@ -66,7 +66,7 @@ class GRApiManager:
"product_user_id": res["product_user_id"],
"email": res["metadata"].get("email_address"),
"display_name": res["metadata"].get("display_name"),
- 'blocked': res['blocked'],
+ "blocked": res["blocked"],
}
)
@@ -122,9 +122,7 @@ class GRApiManager:
AMT account to a General Research account."""
url = f"{self.base_url}/{self.product_id}/user/{amt_worker_id}/"
try:
- self._request(
- "PATCH", url, json={"product_user_id": user.product_user_id}
- )
+ self._request("PATCH", url, json={"product_user_id": user.product_user_id})
except GRApiNotFoundError as exc:
raise ValueError(f"User {amt_worker_id} does not exist") from exc
except GRApiError as exc:
diff --git a/jb/managers/hit.py b/jb/managers/hit.py
index 63af2d4..a178d4d 100644
--- a/jb/managers/hit.py
+++ b/jb/managers/hit.py
@@ -1,11 +1,10 @@
from datetime import datetime, timezone
-from typing import Optional, List
from psycopg import sql
from jb.managers import PostgresManager
from jb.models.definitions import HitStatus
-from jb.models.hit import HitQuestion, HitType, Hit
+from jb.models.hit import Hit, HitQuestion, HitType
class HitQuestionManager(PostgresManager):
@@ -13,21 +12,17 @@ class HitQuestionManager(PostgresManager):
def create(self, question: HitQuestion) -> None:
assert question.id is None
data = question.to_postgres()
- query = sql.SQL(
- """
+ query = sql.SQL("""
INSERT INTO mtwerk_question (url, height)
VALUES (%(url)s, %(height)s)
RETURNING id;
- """
- )
+ """)
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- pk = c.fetchone()["id"] # type: ignore
- conn.commit()
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ pk = c.fetchone()["id"] # type: ignore
+ conn.commit()
question.id = pk
- return None
def get_by_id(self, question_id: int) -> HitQuestion:
res = self.pg_config.execute_sql_query(
@@ -42,7 +37,7 @@ class HitQuestionManager(PostgresManager):
assert len(res) == 1
return HitQuestion.model_validate(res[0])
- def get_by_values_if_exists(self, url: str, height: int) -> Optional[HitQuestion]:
+ def get_by_values_if_exists(self, url: str, height: int) -> HitQuestion | None:
res = self.pg_config.execute_sql_query(
"""
SELECT *
@@ -73,8 +68,7 @@ class HitTypeManager(PostgresManager):
assert hit_type.amt_hit_type_id is not None
data = hit_type.to_postgres()
- query = sql.SQL(
- """
+ query = sql.SQL("""
INSERT INTO mtwerk_hittype (
amt_hit_type_id,
title,
@@ -96,27 +90,21 @@ class HitTypeManager(PostgresManager):
%(min_active)s
)
RETURNING id;
- """
- )
+ """)
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- pk = c.fetchone()["id"] # type: ignore
- conn.commit()
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ pk = c.fetchone()["id"] # type: ignore
+ conn.commit()
hit_type.id = pk
- return None
-
- def filter_active(self) -> List[HitType]:
- res = self.pg_config.execute_sql_query(
- """
+ def filter_active(self) -> list[HitType]:
+ res = self.pg_config.execute_sql_query("""
SELECT *
FROM mtwerk_hittype
WHERE min_active > 0
LIMIT 50
- """
- )
+ """)
if len(res) == 50:
raise ValueError("Too many HitTypes!")
@@ -135,7 +123,7 @@ class HitTypeManager(PostgresManager):
assert len(res) == 1
return HitType.from_postgres(res[0])
- def get_if_exists(self, amt_hit_type_id: str) -> Optional[HitType]:
+ def get_if_exists(self, amt_hit_type_id: str) -> HitType | None:
try:
return self.get(amt_hit_type_id=amt_hit_type_id)
except AssertionError:
@@ -151,22 +139,18 @@ class HitTypeManager(PostgresManager):
def set_min_active(self, hit_type: HitType) -> None:
assert hit_type.id, "must be in the db first!"
- query = sql.SQL(
- """
+ query = sql.SQL("""
UPDATE mtwerk_hittype
SET min_active = %(min_active)s
WHERE id = %(id)s
- """
- )
+ """)
data = {"id": hit_type.id, "min_active": hit_type.min_active}
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- conn.commit()
- row_cnt = c.rowcount
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ conn.commit()
+ row_cnt = c.rowcount
assert row_cnt == 1, f"Expected 1 row updated, got {row_cnt}"
- return None
class HitManager(PostgresManager):
@@ -175,8 +159,7 @@ class HitManager(PostgresManager):
assert hit.amt_hit_id is not None
assert hit.id is None
data = hit.to_postgres()
- query = sql.SQL(
- """
+ query = sql.SQL("""
INSERT INTO mtwerk_hit (
amt_hit_id,
hit_type_id,
@@ -210,38 +193,32 @@ class HitManager(PostgresManager):
%(assignment_available_count)s
)
RETURNING id;
- """
- )
+ """)
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- pk = c.fetchone()["id"] # type: ignore
- conn.commit()
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ pk = c.fetchone()["id"] # type: ignore
+ conn.commit()
hit.id = pk
return hit
def update_status(self, amt_hit_id: str, hit_status: HitStatus):
now = datetime.now(tz=timezone.utc)
- query = sql.SQL(
- """
+ query = sql.SQL("""
UPDATE mtwerk_hit
SET status = %(status)s, modified_at = %(modified_at)s
WHERE amt_hit_id = %(amt_hit_id)s;
- """
- )
+ """)
data = {
"amt_hit_id": amt_hit_id,
"status": hit_status.value,
"modified_at": now,
}
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- conn.commit()
- assert c.rowcount == 1, c.rowcount
- return None
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ conn.commit()
+ assert c.rowcount == 1, c.rowcount
def update_hit(self, hit: Hit):
hit.modified_at = datetime.now(tz=timezone.utc)
@@ -255,8 +232,7 @@ class HitManager(PostgresManager):
"amt_hit_id",
"modified_at",
}
- query = sql.SQL(
- """
+ query = sql.SQL("""
UPDATE mtwerk_hit
SET status = %(status)s, review_status = %(review_status)s,
assignment_pending_count = %(assignment_pending_count)s,
@@ -265,17 +241,14 @@ class HitManager(PostgresManager):
modified_at = %(modified_at)s
WHERE amt_hit_id = %(amt_hit_id)s
RETURNING id;
- """
- )
+ """)
data = hit.model_dump(mode="json", include=fields)
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- conn.commit()
- assert c.rowcount == 1, c.rowcount
- hit.id = c.fetchone()["id"] # type: ignore
- return None
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ conn.commit()
+ assert c.rowcount == 1, c.rowcount
+ hit.id = c.fetchone()["id"] # type: ignore
def get_from_amt_id(self, amt_hit_id: str) -> Hit:
res = self.pg_config.execute_sql_query(
@@ -324,7 +297,7 @@ class HitManager(PostgresManager):
return Hit.from_postgres(res)
- def get_from_amt_id_if_exists(self, amt_hit_id: str) -> Optional[Hit]:
+ def get_from_amt_id_if_exists(self, amt_hit_id: str) -> Hit | None:
try:
return self.get_from_amt_id(amt_hit_id=amt_hit_id)
diff --git a/jb/managers/thl.py b/jb/managers/thl.py
index e50fe76..85e0697 100644
--- a/jb/managers/thl.py
+++ b/jb/managers/thl.py
@@ -1,25 +1,18 @@
+import requests
+from generalresearch.currency import USDCent
+from generalresearch.models.thl.definitions import PayoutStatus
from generalresearch.models.thl.payout import UserPayoutEvent
from generalresearch.models.thl.task_status import TaskStatusResponse
from generalresearch.models.thl.wallet.cashout_method import (
- CashoutRequestResponse,
CashoutRequestInfo,
+ CashoutRequestResponse,
)
-from generalresearch.models.thl.user_profile import UserProfile
-from generalresearch.currency import USDCent
-
from jb.config import settings
-
-from generalresearch.models.thl.definitions import PayoutStatus
-
-
-from typing import Optional
-import requests
-
from jb.models.auth import User
-def get_task_status(tsid: str) -> Optional[TaskStatusResponse]:
+def get_task_status(tsid: str) -> TaskStatusResponse | None:
url = f"{settings.fsb_host}{settings.product_id}/status/{tsid}/"
d = requests.get(url).json()
if d.get("msg") == "invalid tsid":
@@ -67,7 +60,7 @@ def get_wallet_balance(amt_worker_id: str) -> USDCent:
return USDCent(requests.get(url, params=params).json()["wallet"]["amount"])
-def get_wallet_balance_if_non_negative(amt_worker_id: str) -> Optional[USDCent]:
+def get_wallet_balance_if_non_negative(amt_worker_id: str) -> USDCent | None:
url = f"{settings.fsb_host}{settings.product_id}/wallet/"
params = {"bpuid": amt_worker_id}
amt = requests.get(url, params=params).json()["wallet"]["amount"]
diff --git a/jb/managers/worker.py b/jb/managers/worker.py
index e2d7237..9e7d5e3 100644
--- a/jb/managers/worker.py
+++ b/jb/managers/worker.py
@@ -1,5 +1,3 @@
-from typing import List
-
from mypy_boto3_mturk.type_defs import WorkerBlockTypeDef
from jb.decorators import AMT_CLIENT
@@ -8,9 +6,9 @@ from jb.decorators import AMT_CLIENT
class WorkerManager:
@staticmethod
- def fetch_worker_blocks() -> List[WorkerBlockTypeDef]:
+ def fetch_worker_blocks() -> list[WorkerBlockTypeDef]:
p = AMT_CLIENT.get_paginator("list_worker_blocks")
- res: List[WorkerBlockTypeDef] = []
+ res: list[WorkerBlockTypeDef] = []
for item in p.paginate():
res.extend(item["WorkerBlocks"])
return res
diff --git a/jb/models/__init__.py b/jb/models/__init__.py
index 0aeae14..7fe23a7 100644
--- a/jb/models/__init__.py
+++ b/jb/models/__init__.py
@@ -1,7 +1,6 @@
from decimal import Decimal
-from typing import Optional
-from pydantic import BaseModel, Field, ConfigDict
+from pydantic import BaseModel, ConfigDict, Field
class HTTPHeaders(BaseModel):
@@ -12,7 +11,7 @@ class HTTPHeaders(BaseModel):
# 'Mon, 15 Jan 2024 23:40:32 GMT'
date: str = Field()
- connection: Optional[str] = Field(default=None) # 'close'
+ connection: str | None = Field(default=None) # 'close'
class ResponseMetadata(BaseModel):
diff --git a/jb/models/assignment.py b/jb/models/assignment.py
index 92e5a89..775cd63 100644
--- a/jb/models/assignment.py
+++ b/jb/models/assignment.py
@@ -1,17 +1,17 @@
import logging
from datetime import datetime, timezone
-from typing import Optional, TypedDict, Any
+from typing import Any, TypedDict
from xml.etree import ElementTree
from mypy_boto3_mturk.type_defs import AssignmentTypeDef
from pydantic import (
BaseModel,
- Field,
ConfigDict,
- model_validator,
+ Field,
PositiveInt,
TypeAdapter,
ValidationError,
+ model_validator,
)
from typing_extensions import Self
@@ -36,8 +36,8 @@ class AssignmentStub(BaseModel):
validate_assignment=True,
)
- id: Optional[PositiveInt] = Field(default=None)
- hit_id: Optional[PositiveInt] = Field(default=None)
+ id: PositiveInt | None = Field(default=None)
+ hit_id: PositiveInt | None = Field(default=None)
amt_assignment_id: AMTBoto3ID = Field()
amt_hit_id: AMTBoto3ID = Field()
amt_worker_id: str = Field(min_length=3, max_length=50)
@@ -50,7 +50,7 @@ class AssignmentStub(BaseModel):
description="When this record was saved in the database",
)
- modified_at: Optional[AwareDatetimeISO] = Field(
+ modified_at: AwareDatetimeISO | None = Field(
default_factory=lambda: datetime.now(tz=timezone.utc),
description="When this record was updated / modified in the database",
)
@@ -96,18 +96,18 @@ class Assignment(AssignmentStub):
"submitted results.",
)
- approval_time: Optional[AwareDatetimeISO] = Field(
+ approval_time: AwareDatetimeISO | None = Field(
default=None,
description="The date and time the Requester approved the results. This "
"value is omitted from the assignment if the Requester has "
"not yet approved the results.",
)
- rejection_time: Optional[AwareDatetimeISO] = Field(
+ rejection_time: AwareDatetimeISO | None = Field(
default=None,
description="The date and time the Requester rejected the results.",
)
- requester_feedback: Optional[str] = Field(
+ requester_feedback: str | None = Field(
# Default: None. This field isn't returned with assignment data by
# default. To request this field, specify a response group of
# AssignmentFeedback. For information about response groups, see
@@ -123,11 +123,11 @@ class Assignment(AssignmentStub):
},
)
- answer_xml: Optional[str] = Field(default=None, exclude=True)
+ answer_xml: str | None = Field(default=None, exclude=True)
# GRL Specific
- tsid: Optional[UUIDStr] = Field(default=None)
+ tsid: UUIDStr | None = Field(default=None)
# --- Validators ---
@@ -173,7 +173,7 @@ class Assignment(AssignmentStub):
# --- Properties ---
@property
- def answers_dict(self) -> Optional[AnswerDict]:
+ def answers_dict(self) -> AnswerDict | None:
# See https://docs.aws.amazon.com/AWSMechTurk/latest/AWSMturkAPI/ApiReference_AssignmentDataStructureArticle.html
# https://docs.aws.amazon.com/AWSMechTurk/latest/AWSMechanicalTurkRequester/Concepts_NotificationsArticle.html
if self.answer_xml is None:
diff --git a/jb/models/bonus.py b/jb/models/bonus.py
index 5f81add..2c1d00c 100644
--- a/jb/models/bonus.py
+++ b/jb/models/bonus.py
@@ -1,9 +1,9 @@
-from typing import Optional, Dict, Any
+from typing import Any
-from pydantic import BaseModel, Field, ConfigDict, PositiveInt
+from generalresearch.currency import USDCent
+from pydantic import BaseModel, ConfigDict, Field, PositiveInt
from typing_extensions import Self
-from generalresearch.currency import USDCent
from jb.models.custom_types import AMTBoto3ID, AwareDatetimeISO, UUIDStr
@@ -20,8 +20,8 @@ class Bonus(BaseModel):
extra="forbid",
validate_assignment=True,
)
- id: Optional[PositiveInt] = Field(default=None)
- assignment_id: Optional[PositiveInt] = Field(default=None)
+ id: PositiveInt | None = Field(default=None)
+ assignment_id: PositiveInt | None = Field(default=None)
amt_worker_id: str = Field(min_length=3, max_length=50)
amt_assignment_id: AMTBoto3ID = Field()
@@ -40,7 +40,7 @@ class Bonus(BaseModel):
return d
@classmethod
- def from_postgres(cls, data: Dict[str, Any]) -> Self:
+ def from_postgres(cls, data: dict[str, Any]) -> Self:
data["amount"] = USDCent(round(data["amount"] * 100))
fields = set(cls.model_fields.keys())
data = {k: v for k, v in data.items() if k in fields}
diff --git a/jb/models/custom_types.py b/jb/models/custom_types.py
index 10bc9d1..a58dcb7 100644
--- a/jb/models/custom_types.py
+++ b/jb/models/custom_types.py
@@ -1,19 +1,18 @@
import re
from datetime import datetime, timezone
-from typing import Any, Optional
+from typing import Annotated, Any
from uuid import UUID
from pydantic import (
AwareDatetime,
+ HttpUrl,
StringConstraints,
TypeAdapter,
- HttpUrl,
)
from pydantic.functional_serializers import PlainSerializer
from pydantic.functional_validators import AfterValidator, BeforeValidator
from pydantic.networks import UrlConstraints
from pydantic_core import Url
-from typing_extensions import Annotated
def convert_datetime_to_iso_8601_with_z_suffix(dt: datetime) -> str:
@@ -22,7 +21,7 @@ def convert_datetime_to_iso_8601_with_z_suffix(dt: datetime) -> str:
return dt.strftime("%Y-%m-%dT%H:%M:%S.%fZ")
-def convert_str_dt(v: Any) -> Optional[AwareDatetime]:
+def convert_str_dt(v: Any) -> AwareDatetime | None:
# By default, pydantic is unable to handle tz-aware isoformat str. Attempt to parse a str
# that was dumped using the iso8601 format with Z suffix.
if v is not None and type(v) is str:
diff --git a/jb/models/errors.py b/jb/models/errors.py
index 94f5fbb..c590c6a 100644
--- a/jb/models/errors.py
+++ b/jb/models/errors.py
@@ -1,7 +1,7 @@
import re
from enum import Enum
-from pydantic import BaseModel, Field, ConfigDict, model_validator
+from pydantic import BaseModel, ConfigDict, Field, model_validator
from jb.models import ResponseMetadata
diff --git a/jb/models/event.py b/jb/models/event.py
index f8867c0..0016ca7 100644
--- a/jb/models/event.py
+++ b/jb/models/event.py
@@ -1,9 +1,9 @@
-from typing import Dict, Any
+from typing import Any
from mypy_boto3_mturk.literals import EventTypeType
from pydantic import BaseModel, Field
-from jb.models.custom_types import AwareDatetimeISO, AMTBoto3ID
+from jb.models.custom_types import AMTBoto3ID, AwareDatetimeISO
class MTurkEvent(BaseModel):
@@ -29,7 +29,7 @@ class MTurkEvent(BaseModel):
)
@classmethod
- def from_sns(cls, data: Dict[str, Any]):
+ def from_sns(cls, data: dict[str, Any]):
return cls.model_validate(
{
"event_type": data["EventType"],
diff --git a/jb/models/hit.py b/jb/models/hit.py
index 45478fc..f6c854f 100644
--- a/jb/models/hit.py
+++ b/jb/models/hit.py
@@ -1,25 +1,25 @@
-from datetime import datetime, timezone, timedelta
-from typing import Optional, List, Dict, Any
+from datetime import datetime, timedelta, timezone
+from typing import Any
from uuid import uuid4
from xml.etree import ElementTree
+from generalresearch.currency import USDCent
from mypy_boto3_mturk.type_defs import HITTypeDef
from pydantic import (
BaseModel,
- Field,
- PositiveInt,
ConfigDict,
+ Field,
NonNegativeInt,
+ PositiveInt,
)
from typing_extensions import Self
-from generalresearch.currency import USDCent
-from jb.models.custom_types import AMTBoto3ID, HttpsUrlStr, AwareDatetimeISO
-from jb.models.definitions import HitStatus, HitReviewStatus
+from jb.models.custom_types import AMTBoto3ID, AwareDatetimeISO, HttpsUrlStr
+from jb.models.definitions import HitReviewStatus, HitStatus
class HitQuestion(BaseModel):
- id: Optional[PositiveInt] = Field(default=None)
+ id: PositiveInt | None = Field(default=None)
url: HttpsUrlStr = Field()
height: PositiveInt = Field(default=1_200, ge=100, le=4_000)
@@ -33,7 +33,7 @@ class HitQuestion(BaseModel):
def xml(self) -> str:
return f"""<?xml version="1.0" encoding="UTF-8"?>
<ExternalQuestion xmlns="http://mechanicalturk.amazonaws.com/AWSMechanicalTurkDataSchemas/2006-07-14/ExternalQuestion.xsd">
- <ExternalURL>{str(self.url)}</ExternalURL>
+ <ExternalURL>{self.url!s}</ExternalURL>
<FrameHeight>{self.height}</FrameHeight>
</ExternalQuestion>"""
@@ -82,8 +82,8 @@ class HitType(HitTypeCommon):
https://docs.aws.amazon.com/AWSMechTurk/latest/AWSMturkAPI/ApiReference_CreateHITTypeOperation.html
"""
- id: Optional[PositiveInt] = Field(default=None)
- amt_hit_type_id: Optional[AMTBoto3ID] = Field(default=None)
+ id: PositiveInt | None = Field(default=None)
+ amt_hit_type_id: AMTBoto3ID | None = Field(default=None)
# --- GRL Specific ---
min_active: NonNegativeInt = Field(default=0, le=100_000)
@@ -104,11 +104,11 @@ class HitType(HitTypeCommon):
return d
@classmethod
- def from_postgres(cls, data: Dict[str, Any]) -> Self:
+ def from_postgres(cls, data: dict[str, Any]) -> Self:
data["reward"] = USDCent(round(data["reward"] * 100))
return cls.model_validate(data)
- def generate_hit_amt_request(self, question: HitQuestion) -> Dict[str, Any]:
+ def generate_hit_amt_request(self, question: HitQuestion) -> dict[str, Any]:
d = dict()
d["HITTypeId"] = self.amt_hit_type_id
d["MaxAssignments"] = 1
@@ -124,9 +124,9 @@ class Hit(HitTypeCommon):
validate_assignment=True,
)
- id: Optional[PositiveInt] = Field(default=None)
- hit_type_id: Optional[PositiveInt] = Field(default=None)
- question_id: Optional[PositiveInt] = Field(default=None)
+ id: PositiveInt | None = Field(default=None)
+ hit_type_id: PositiveInt | None = Field(default=None)
+ question_id: PositiveInt | None = Field(default=None)
amt_hit_id: AMTBoto3ID = Field()
amt_hit_type_id: AMTBoto3ID = Field()
@@ -138,10 +138,8 @@ class Hit(HitTypeCommon):
# TODO: Check if this is actually ever going to be None. I type fixed it,
# but I don't have anything to suggest it isn't requred. -- Max 2026-02-24
- creation_time: Optional[AwareDatetimeISO] = Field(
- default=None, description="From aws"
- )
- expiration: Optional[AwareDatetimeISO] = Field(default=None)
+ creation_time: AwareDatetimeISO | None = Field(default=None, description="From aws")
+ expiration: AwareDatetimeISO | None = Field(default=None)
# GRL Specific
created_at: AwareDatetimeISO = Field(
@@ -155,7 +153,7 @@ class Hit(HitTypeCommon):
# -- Hit specific
- qualification_requirements: Optional[List[Dict[str, Any]]] = Field(default=None)
+ qualification_requirements: list[dict[str, Any]] | None = Field(default=None)
max_assignments: int = Field()
# # this comes back as expiration. only for the request
@@ -235,7 +233,7 @@ class Hit(HitTypeCommon):
return d
@classmethod
- def from_postgres(cls, data: Dict[str, Any]) -> Self:
+ def from_postgres(cls, data: dict[str, Any]) -> Self:
data["reward"] = USDCent(round(data["reward"] * 100))
return cls.model_validate(data)
diff --git a/jb/settings.py b/jb/settings.py
index 86c8a36..7747afc 100644
--- a/jb/settings.py
+++ b/jb/settings.py
@@ -1,10 +1,9 @@
import os
from functools import lru_cache
from pathlib import Path
-from typing import Optional
from generalresearch.models.custom_types import InfluxDsn
-from pydantic import Field, PostgresDsn, HttpUrl, RedisDsn, SecretStr
+from pydantic import Field, HttpUrl, PostgresDsn, RedisDsn, SecretStr
from pydantic_settings import BaseSettings, SettingsConfigDict
from jb.models.custom_types import UUIDStr
@@ -18,14 +17,14 @@ BASE_HTML = BASE_HTML_PATH.read_text()
class AmtJbBaseSettings(BaseSettings):
debug: bool = Field(default=True)
- redis: Optional[RedisDsn] = Field(default=None)
+ redis: RedisDsn | None = Field(default=None)
redis_timeout: float = Field(default=0.10)
amt_jb_db: PostgresDsn = Field()
- amt_endpoint: Optional[HttpUrl] = Field(default=None)
- amt_access_id: Optional[str] = Field(default=None)
- amt_secret_key: Optional[str] = Field(default=None)
+ amt_endpoint: HttpUrl | None = Field(default=None)
+ amt_access_id: str | None = Field(default=None)
+ amt_secret_key: str | None = Field(default=None)
aws_owner_id: str = Field()
aws_subscription_arn: str = Field()
@@ -45,11 +44,11 @@ class Settings(AmtJbBaseSettings):
fsb_host: HttpUrl = Field(default=HttpUrl("https://fsb.generalresearch.com/"))
# Needed for admin function on fsb w/o authentication
- fsb_host_private_route: Optional[str] = Field(default=None)
+ fsb_host_private_route: str | None = Field(default=None)
product_id: UUIDStr = Field()
- influx_db: Optional[InfluxDsn] = Field(default=None)
+ influx_db: InfluxDsn | None = Field(default=None)
sns_path: str = Field()
@@ -58,9 +57,7 @@ class Settings(AmtJbBaseSettings):
magic_token_salt: SecretStr = Field(min_length=32)
- gr_api_host: HttpUrl = Field(
- default=HttpUrl("https://generalresearch.com/api/v2/")
- )
+ gr_api_host: HttpUrl = Field(default=HttpUrl("https://generalresearch.com/api/v2/"))
gr_api_token: SecretStr = Field(min_length=1)
mautic_api_key: SecretStr = Field(min_length=32)