Coverage for polar/worker/_enqueue.py: 96%
86 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 12:42 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 12:42 +0000
1import contextlib
2import contextvars
3import itertools
4import uuid
5from collections import defaultdict
6from collections.abc import AsyncIterator, Iterable, Mapping
7from typing import Any, Self, TypeAlias
9import dramatiq
10import structlog
12from polar.logging import Logger
13from polar.redis import Redis
15log: Logger = structlog.get_logger()
18JSONSerializable: TypeAlias = (
19 Mapping[str, "JSONSerializable"]
20 | Iterable["JSONSerializable"]
21 | str
22 | int
23 | float
24 | bool
25 | uuid.UUID
26 | None
27)
30_job_queue_manager: contextvars.ContextVar["JobQueueManager | None"] = (
31 contextvars.ContextVar("polar.job_queue_manager")
32)
34FLUSH_BATCH_SIZE = 50
37class JobQueueManager:
38 __slots__ = ("_enqueued_jobs", "_ingested_events")
40 def __init__(self) -> None:
41 self._enqueued_jobs: list[
42 tuple[str, tuple[JSONSerializable, ...], dict[str, JSONSerializable]]
43 ] = []
44 self._ingested_events: list[uuid.UUID] = []
46 def enqueue_job(
47 self, actor: str, *args: JSONSerializable, **kwargs: JSONSerializable
48 ) -> None:
49 self._enqueued_jobs.append((actor, args, kwargs))
50 log.debug("polar.worker.job_enqueued", actor=actor)
52 def enqueue_events(self, *event_ids: uuid.UUID) -> None:
53 self._ingested_events.extend(event_ids)
55 async def flush(self, broker: dramatiq.Broker, redis: Redis) -> None:
56 if len(self._ingested_events) > 0: 56 ↛ 57line 56 didn't jump to line 57 because the condition on line 56 was never true
57 self.enqueue_job("event.ingested", self._ingested_events)
59 if not self._enqueued_jobs:
60 self.reset()
61 return
63 queue_messages = defaultdict[str, list[tuple[str, Any]]](list)
64 all_messages: list[tuple[str, Any]] = []
66 for actor_name, args, kwargs in self._enqueued_jobs:
67 fn: dramatiq.Actor[Any, Any] = broker.get_actor(actor_name)
68 redis_message_id = str(uuid.uuid4())
69 message = fn.message_with_options(
70 args=args, kwargs=kwargs, redis_message_id=redis_message_id
71 )
72 encoded_message = message.encode()
73 queue_messages[message.queue_name].append(
74 (redis_message_id, encoded_message)
75 )
76 all_messages.append((fn.actor_name, message.encode()))
78 for queue_name, messages in queue_messages.items():
79 for batch in itertools.batched(messages, FLUSH_BATCH_SIZE):
80 await self._batch_hset_messages(redis, queue_name, batch)
81 await self._batch_rpush_queue(
82 redis, queue_name, (message_id for message_id, _ in batch)
83 )
85 for actor_name, encoded_message in all_messages:
86 log.debug(
87 "polar.worker.job_flushed", actor=actor_name, message=encoded_message
88 )
90 self.reset()
92 async def _batch_hset_messages(
93 self,
94 redis: Redis,
95 queue_name: str,
96 message_batch: Iterable[tuple[str, Any]],
97 ) -> None:
98 """Batch hset operations for message storage."""
99 hash_key = f"dramatiq:{queue_name}.msgs"
100 await redis.hset(
101 hash_key,
102 mapping={
103 message_id: encoded_message
104 for message_id, encoded_message in message_batch
105 },
106 )
108 async def _batch_rpush_queue(
109 self, redis: Redis, queue_name: str, message_ids: Iterable[str]
110 ) -> None:
111 """Batch rpush operations for queue entries."""
112 queue_key = f"dramatiq:{queue_name}"
113 await redis.rpush(queue_key, *message_ids)
115 def reset(self) -> None:
116 self._enqueued_jobs = []
117 self._ingested_events = []
119 @classmethod
120 def set(cls) -> "Self":
121 job_queue_manager = cls()
122 _job_queue_manager.set(job_queue_manager)
123 return job_queue_manager
125 @classmethod
126 def close(cls) -> None:
127 job_queue_manager = cls.get()
128 job_queue_manager.reset()
129 _job_queue_manager.set(None)
131 @classmethod
132 @contextlib.asynccontextmanager
133 async def open(cls, broker: dramatiq.Broker, redis: Redis) -> AsyncIterator["Self"]:
134 job_queue_manager = cls.set()
135 try:
136 yield job_queue_manager
137 await job_queue_manager.flush(broker, redis)
138 finally:
139 cls.close()
141 @classmethod
142 def get(cls) -> "JobQueueManager":
143 job_queue_manager = _job_queue_manager.get()
144 if job_queue_manager is None: 144 ↛ 145line 144 didn't jump to line 145 because the condition on line 144 was never true
145 raise RuntimeError("JobQueueManager not initialized")
146 return job_queue_manager
149def enqueue_job(
150 actor: str, *args: JSONSerializable, **kwargs: JSONSerializable
151) -> None:
152 """Enqueue a job by actor name."""
153 job_queue_manager = JobQueueManager.get()
154 job_queue_manager.enqueue_job(actor, *args, **kwargs)
157def enqueue_events(*event_ids: uuid.UUID) -> None:
158 """Enqueue events to be ingested."""
159 job_queue_manager = JobQueueManager.get()
160 job_queue_manager.enqueue_events(*event_ids)