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

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 

8 

9import dramatiq 

10import structlog 

11 

12from polar.logging import Logger 

13from polar.redis import Redis 

14 

15log: Logger = structlog.get_logger() 

16 

17 

18JSONSerializable: TypeAlias = ( 

19 Mapping[str, "JSONSerializable"] 

20 | Iterable["JSONSerializable"] 

21 | str 

22 | int 

23 | float 

24 | bool 

25 | uuid.UUID 

26 | None 

27) 

28 

29 

30_job_queue_manager: contextvars.ContextVar["JobQueueManager | None"] = ( 

31 contextvars.ContextVar("polar.job_queue_manager") 

32) 

33 

34FLUSH_BATCH_SIZE = 50 

35 

36 

37class JobQueueManager: 

38 __slots__ = ("_enqueued_jobs", "_ingested_events") 

39 

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] = [] 

45 

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) 

51 

52 def enqueue_events(self, *event_ids: uuid.UUID) -> None: 

53 self._ingested_events.extend(event_ids) 

54 

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) 

58 

59 if not self._enqueued_jobs: 

60 self.reset() 

61 return 

62 

63 queue_messages = defaultdict[str, list[tuple[str, Any]]](list) 

64 all_messages: list[tuple[str, Any]] = [] 

65 

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())) 

77 

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 ) 

84 

85 for actor_name, encoded_message in all_messages: 

86 log.debug( 

87 "polar.worker.job_flushed", actor=actor_name, message=encoded_message 

88 ) 

89 

90 self.reset() 

91 

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 ) 

107 

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) 

114 

115 def reset(self) -> None: 

116 self._enqueued_jobs = [] 

117 self._ingested_events = [] 

118 

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 

124 

125 @classmethod 

126 def close(cls) -> None: 

127 job_queue_manager = cls.get() 

128 job_queue_manager.reset() 

129 _job_queue_manager.set(None) 

130 

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() 

140 

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 

147 

148 

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) 

155 

156 

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)