aboutsummaryrefslogtreecommitdiff
path: root/jb/views/common.py
diff options
context:
space:
mode:
Diffstat (limited to 'jb/views/common.py')
-rw-r--r--jb/views/common.py18
1 files changed, 13 insertions, 5 deletions
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)