aboutsummaryrefslogtreecommitdiff
path: root/jb
diff options
context:
space:
mode:
authorMax Nanis2026-09-10 16:59:43 -0700
committerMax Nanis2026-09-10 16:59:43 -0700
commit0ac38066402762dc5fa24f2ffa2170681a6eddb1 (patch)
tree923ea3f9ccd7adaaac89d20074d7ed6af532260e /jb
parentcb098c80351727055ce982b9c2631aab7277b772 (diff)
downloadamt-jb-0ac38066402762dc5fa24f2ffa2170681a6eddb1.tar.gz
amt-jb-0ac38066402762dc5fa24f2ffa2170681a6eddb1.zip
http/managers/models green.
Diffstat (limited to 'jb')
-rw-r--r--jb/api/magic_token.py16
-rw-r--r--jb/decorators.py10
-rw-r--r--jb/flow/events.py16
-rw-r--r--jb/views/common.py18
4 files changed, 44 insertions, 16 deletions
diff --git a/jb/api/magic_token.py b/jb/api/magic_token.py
index 5f78996..7b6c1aa 100644
--- a/jb/api/magic_token.py
+++ b/jb/api/magic_token.py
@@ -3,7 +3,7 @@ import secrets
from fastapi import HTTPException, status
-from jb.decorators import REDIS
+from jb.decorators import get_redis
from jb.models.auth import AmtAccountLink
MAGIC_TOKEN_PREFIX = "auth:magic:"
@@ -25,7 +25,8 @@ def create_magic_token(user_email: str) -> str:
raise ValueError("user_email must not be empty")
token = secrets.token_urlsafe(32)
- REDIS.set(
+ redis_client = get_redis()
+ redis_client.set(
redis_token_key(token),
user_email,
ex=MAGIC_TOKEN_TTL,
@@ -34,7 +35,8 @@ def create_magic_token(user_email: str) -> str:
def consume_magic_token(token: str) -> str:
- user_email = REDIS.getdel(redis_token_key(token))
+ redis_client = get_redis()
+ user_email = redis_client.getdel(redis_token_key(token))
if user_email is None:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
@@ -45,12 +47,13 @@ def consume_magic_token(token: str) -> str:
def create_amt_account_link_token(email: str, amt_worker_id: str) -> str:
"""Bind an email and AMT worker ID to an opaque, short-lived token."""
+ redis_client = get_redis()
data = AmtAccountLink(
email=email,
amt_worker_id=amt_worker_id,
)
token = secrets.token_urlsafe(32)
- REDIS.set(
+ redis_client.set(
redis_token_key(token, AMT_ACCOUNT_LINK_TOKEN_PREFIX),
data.model_dump_json(),
ex=MAGIC_TOKEN_TTL,
@@ -60,7 +63,10 @@ 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))
+ redis_client = get_redis()
+ raw_data = redis_client.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/decorators.py b/jb/decorators.py
index 6e1336d..9c7a31c 100644
--- a/jb/decorators.py
+++ b/jb/decorators.py
@@ -22,7 +22,15 @@ redis_config = RedisConfig(
socket_timeout=settings.redis_timeout,
socket_connect_timeout=settings.redis_timeout,
)
-REDIS = redis_config.create_redis_client()
+
+
+def get_redis_config():
+ return redis_config
+
+
+def get_redis():
+ return redis_config.create_redis_client()
+
# --- Logging ---
diff --git a/jb/flow/events.py b/jb/flow/events.py
index f224c01..bee63fd 100644
--- a/jb/flow/events.py
+++ b/jb/flow/events.py
@@ -10,7 +10,7 @@ from jb.config import (
CONSUMER_NAME,
JB_EVENTS_STREAM,
)
-from jb.decorators import LOG, REDIS
+from jb.decorators import LOG, get_redis
from jb.flow.assignment_tasks import process_assignment_submitted
from jb.flow.monitoring import emit_error_event
from jb.models.event import MTurkEvent
@@ -39,7 +39,10 @@ def process_mturk_events(executor: Executor):
def create_consumer_group():
try:
- REDIS.xgroup_create(JB_EVENTS_STREAM, CONSUMER_GROUP, id="0", mkstream=True)
+ redis_client = get_redis()
+ redis_client.xgroup_create(
+ JB_EVENTS_STREAM, CONSUMER_GROUP, id="0", mkstream=True
+ )
except redis.exceptions.ResponseError as e:
if "BUSYGROUP Consumer Group name already exists" in str(e):
pass # group already exists
@@ -48,7 +51,8 @@ def create_consumer_group():
def process_mturk_events_chunk(executor: Executor) -> int | None:
- msgs_raw = REDIS.xreadgroup(
+ redis_client = get_redis()
+ msgs_raw = redis_client.xreadgroup(
groupname=CONSUMER_GROUP,
consumername=CONSUMER_NAME,
streams={JB_EVENTS_STREAM: ">"},
@@ -71,7 +75,7 @@ def process_mturk_events_chunk(executor: Executor) -> int | None:
)
else:
LOG.info(f"Discarding {event}")
- REDIS.xdel(JB_EVENTS_STREAM, msg_id)
+ redis_client.xdel(JB_EVENTS_STREAM, msg_id)
futures.wait(fs, timeout=60)
return len(msgs)
@@ -80,6 +84,8 @@ def process_mturk_events_chunk(executor: Executor) -> int | None:
def process_assignment_submitted_event(event: MTurkEvent, msg_id: str):
from jb.decorators import AM, AMTM, BM, HM
+ redis_client = get_redis()
+
try:
process_assignment_submitted(amtm=AMTM, am=AM, hm=HM, bm=BM, event=event)
except Exception as e:
@@ -89,4 +95,4 @@ def process_assignment_submitted_event(event: MTurkEvent, msg_id: str):
amt_hit_type_id=event.amt_hit_type_id,
)
- REDIS.xackdel(JB_EVENTS_STREAM, CONSUMER_GROUP, msg_id)
+ redis_client.xackdel(JB_EVENTS_STREAM, CONSUMER_GROUP, msg_id)
diff --git a/jb/views/common.py b/jb/views/common.py
index d3b3e93..9eee453 100644
--- a/jb/views/common.py
+++ b/jb/views/common.py
@@ -4,11 +4,12 @@ from typing import Annotated, Any
import requests
from fastapi import APIRouter, Depends, HTTPException, Request
from fastapi.responses import HTMLResponse
+from generalresearch.redis_helper import RedisConfig
from starlette.responses import RedirectResponse
from jb.api.auth import get_authenticated_user
from jb.config import JB_EVENTS_STREAM, settings
-from jb.decorators import REDIS
+from jb.decorators import get_redis_config
from jb.flow.monitoring import emit_mturk_notification_event
from jb.models.auth import User
from jb.models.event import MTurkEvent
@@ -41,8 +42,11 @@ async def work(request: Request):
return HTMLResponse(BASE_HTML)
+RedisConfigDep = Annotated[RedisConfig, Depends(get_redis_config)]
+
+
@common_router.post(path=f"/{settings.sns_path}/", include_in_schema=False)
-async def mturk_notifications(request: Request):
+async def mturk_notifications(request: Request, redis_config: RedisConfigDep):
"""
Our SNS topic will POST to this endpoint whenever we get a new message
"""
@@ -60,7 +64,7 @@ async def mturk_notifications(request: Request):
case "Notification":
msg = json.loads(message["Message"])
print("Received MTurk event:", msg)
- enqueue_mturk_notifications(msg)
+ enqueue_mturk_notifications(msg=msg, redis_config=redis_config)
case _:
raise HTTPException(status_code=500, detail="Invalid JSON")
@@ -68,13 +72,17 @@ async def mturk_notifications(request: Request):
return {"status": "ok"}
-def enqueue_mturk_notifications(msg: dict[str, Any]) -> None:
+def enqueue_mturk_notifications(msg: dict[str, Any], redis_config: RedisConfig) -> None:
+ redis_client = redis_config.create_redis_client()
+
for evt in msg["Events"]:
event = MTurkEvent.from_sns(evt)
emit_mturk_notification_event(
event_type=event.event_type, amt_hit_type_id=event.amt_hit_type_id
)
- REDIS.xadd(JB_EVENTS_STREAM, {"data": event.model_dump_json()})
+
+ print("enqueue_mturk_notifications", event.model_dump_json())
+ redis_client.xadd(JB_EVENTS_STREAM, {"data": event.model_dump_json()})
@common_router.get(path="/work/direct/", response_class=HTMLResponse)