Coverage for /home/admin/Documents/AI/applications/lexigram-dev/lexigram/experimental/ai/lexigram-ai-workers/src/lexigram/ai/workers/batch_embedding/worker.py: 21%

135 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-08-25 07:19 +0800

1""" 

2Main batch embedding worker implementation. 

3""" 

4 

5from __future__ import annotations 

6 

7import asyncio 

8from dataclasses import dataclass 

9from datetime import UTC, datetime 

10from typing import TYPE_CHECKING, Any, cast 

11 

12from lexigram.ai.workers.batch_embedding.cache import EmbeddingCache 

13from lexigram.ai.workers.batch_embedding.progress import ProgressTracker 

14from lexigram.ai.workers.batch_embedding.types import ( 

15 BatchEmbeddingProgress, 

16 BatchEmbeddingResult, 

17 EmbeddingStatus, 

18) 

19from lexigram.concurrency import Parallel 

20from lexigram.contracts.core import ExecutionStrategy 

21from lexigram.contracts.core.health import HealthCheckResult, HealthStatus 

22from lexigram.logging import ( 

23 get_logger, 

24) 

25 

26if TYPE_CHECKING: 

27 from lexigram.ai.workers.batch_embedding.protocols import EmbeddingProvider 

28 from lexigram.contracts import VectorStoreProtocol 

29 from lexigram.contracts.ai.rag import ChunkProtocol 

30 from lexigram.contracts.infra.tasks import TaskQueueProtocol 

31 from lexigram.contracts.infra.tasks import TaskWorkerProtocol as TaskWorker 

32 

33logger = get_logger(__name__) 

34 

35 

36@dataclass(slots=True) 

37class _EmbeddingChunk: 

38 """Internal normalized chunk shape for batch embedding jobs.""" 

39 

40 text: str 

41 metadata: dict[str, Any] | None = None 

42 

43 

44class BatchEmbeddingWorker: 

45 """ 

46 Background worker for batch embedding generation. 

47 

48 Optimizes embedding generation through: 

49 - Batch processing (minimize API calls) 

50 - Cache integration (avoid redundant computations) 

51 - Progress tracking (resumable on failure) 

52 - Concurrent processing (parallel batches) 

53 

54 Example: 

55 ```python 

56 from lexigram.ai.workers.batch_embedding import BatchEmbeddingWorker 

57 from lexigram.contracts import VectorStoreProtocol 

58 from lexigram.contracts.infra.tasks import TaskQueueProtocol 

59 

60 # Setup — resolve implementations via the DI container 

61 vector_store = container.resolve(VectorStoreProtocol) 

62 embedding_provider = container.resolve(EmbeddingClientProtocol) 

63 queue = container.resolve(TaskQueueProtocol) 

64 

65 worker = BatchEmbeddingWorker( 

66 vector_store=vector_store, 

67 embedding_provider=embedding_provider, 

68 queue=queue, 

69 concurrency=3, 

70 ) 

71 

72 # Start worker 

73 await worker.start() 

74 

75 # Submit embedding job 

76 job_id = await worker.embed_batch( 

77 chunks=chunks, 

78 collection_name="my-docs", 

79 model_name="text-embedding-ada-002", 

80 batch_size=100, 

81 ) 

82 

83 # Check progress 

84 progress = await worker.get_progress(job_id) 

85 logger.info(f"Progress: {progress.progress_percent:.1f}%") 

86 logger.info(f"Cache hit rate: {progress.cache_hit_rate:.1f}%") 

87 

88 # Stop worker 

89 await worker.stop() 

90 ``` 

91 """ 

92 

93 def __init__( 

94 self, 

95 vector_store: VectorStoreProtocol, 

96 embedding_provider: EmbeddingProvider, 

97 queue: TaskQueueProtocol, 

98 worker_id: str = "batch-embedding", 

99 concurrency: int = 3, 

100 default_batch_size: int = 100, 

101 enable_cache: bool = True, 

102 ): 

103 """ 

104 Initialize batch embedding worker. 

105 

106 Args: 

107 vector_store: Vector store for storing embeddings 

108 embedding_provider: Provider for generating embeddings 

109 queue: Task queue for job management 

110 worker_id: Unique worker identifier 

111 concurrency: Number of concurrent batch processing tasks 

112 default_batch_size: Default texts per batch 

113 enable_cache: Enable embedding cache 

114 """ 

115 self.vector_store = vector_store 

116 self.embedding_provider = embedding_provider 

117 self.queue = queue 

118 self.worker_id = worker_id 

119 self.concurrency = concurrency 

120 self.default_batch_size = default_batch_size 

121 self.enable_cache = enable_cache 

122 

123 # Component initialization 

124 self._progress_tracker = ProgressTracker() 

125 self._cache = EmbeddingCache() if enable_cache else None 

126 

127 # Worker setup 

128 self._workers: list[TaskWorker] = [] 

129 self._running = False 

130 

131 async def start(self) -> None: 

132 """Start the embedding worker pool.""" 

133 if self._running: 

134 logger.warning("Worker %s already running", self.worker_id) 

135 return 

136 

137 self._running = True 

138 

139 # Create handler registry 

140 handlers = { 

141 "batch_embed": self._handle_batch_embed, 

142 } 

143 

144 # Dynamically resolve concrete worker implementation via DI 

145 from lexigram.contracts.exceptions import ( 

146 DependencyError, 

147 UnresolvableDependencyError, 

148 ) 

149 from lexigram.di.resolution.context import get_resolver 

150 

151 try: 

152 # Try to resolve a factory or class 

153 resolver = get_resolver(self) 

154 if resolver is None: 

155 raise ValueError("No resolver found in context") 

156 worker_class: type[TaskWorker] = await cast("Any", resolver).resolve( 

157 "TaskWorkerClass" 

158 ) 

159 except ( 

160 DependencyError, 

161 UnresolvableDependencyError, 

162 KeyError, 

163 RuntimeError, 

164 ValueError, 

165 ): 

166 # TaskWorker could not be resolved from the DI container. 

167 # Without a concrete worker implementation the pool cannot start. 

168 logger.warning( 

169 "TaskWorker not resolvable from DI container; " 

170 "batch-embedding worker pool will not start. " 

171 "Register a TaskWorkerProtocol implementation before booting." 

172 ) 

173 self._running = False 

174 return 

175 

176 # Create workers 

177 for i in range(self.concurrency): 

178 worker = worker_class( 

179 worker_id=f"{self.worker_id}-{i}", 

180 queue=self.queue, 

181 handler_registry=handlers, 

182 ) 

183 self._workers.append(worker) 

184 

185 # Start all workers concurrently 

186 start_tasks = [worker.start() for worker in self._workers] 

187 await Parallel.execute(*start_tasks, strategy=ExecutionStrategy.ALL_SETTLED) 

188 

189 logger.info( 

190 "Started batch embedding worker pool", 

191 worker_id=self.worker_id, 

192 concurrency=self.concurrency, 

193 ) 

194 

195 async def stop(self) -> None: 

196 """Stop the embedding worker pool.""" 

197 if not self._running: 

198 return 

199 

200 self._running = False 

201 

202 # Stop all workers concurrently 

203 stop_tasks = [worker.stop() for worker in self._workers] 

204 await Parallel.execute(*stop_tasks, strategy=ExecutionStrategy.ALL_SETTLED) 

205 

206 self._workers.clear() 

207 

208 logger.info("Stopped batch embedding worker pool", worker_id=self.worker_id) 

209 

210 async def embed_batch( 

211 self, 

212 chunks: list[ChunkProtocol], 

213 collection_name: str, 

214 model_name: str = "text-embedding-ada-002", 

215 batch_size: int | None = None, 

216 use_cache: bool = True, 

217 priority: int = 0, 

218 ) -> str: 

219 """ 

220 Submit batch embedding job. 

221 

222 Args: 

223 chunks: List of text chunks to embed 

224 collection_name: Vector store collection name 

225 model_name: Embedding model name 

226 batch_size: Texts per batch (default: worker default) 

227 use_cache: Use embedding cache 

228 priority: JobProtocol priority (higher = sooner) 

229 

230 Returns: 

231 JobProtocol ID for tracking 

232 """ 

233 if batch_size is None: 

234 batch_size = self.default_batch_size 

235 

236 # Create job data 

237 job_data = { 

238 "chunks": [{"text": c.text, "metadata": c.metadata or {}} for c in chunks], 

239 "collection_name": collection_name, 

240 "model_name": model_name, 

241 "batch_size": batch_size, 

242 "use_cache": use_cache and self.enable_cache, 

243 } 

244 

245 # Enqueue job 

246 enqueue_result = await self.queue.enqueue( 

247 { 

248 "name": "batch_embed", 

249 "args": (), 

250 "kwargs": job_data, 

251 "priority": priority, 

252 } 

253 ) 

254 if enqueue_result.is_err(): 

255 msg = f"Failed to enqueue embedding job: {enqueue_result.unwrap_err()}" 

256 raise RuntimeError(msg) 

257 

258 job_id = enqueue_result.unwrap() 

259 

260 # Initialize progress tracking 

261 await self._progress_tracker.initialize_job( 

262 job_id=job_id, 

263 total_texts=len(chunks), 

264 ) 

265 

266 logger.info( 

267 "Submitted batch embedding job", 

268 job_id=job_id, 

269 total_chunks=len(chunks), 

270 batch_size=batch_size, 

271 model=model_name, 

272 ) 

273 

274 return job_id 

275 

276 async def get_progress(self, job_id: str) -> BatchEmbeddingProgress | None: 

277 """Get embedding progress for job.""" 

278 return await self._progress_tracker.get_progress(job_id) 

279 

280 async def _handle_batch_embed( 

281 self, 

282 chunks: list[dict[str, Any]], 

283 collection_name: str, 

284 model_name: str, 

285 batch_size: int, 

286 use_cache: bool, 

287 ) -> BatchEmbeddingResult: 

288 """ 

289 Handle batch embedding job. 

290 

291 This is the main worker function that processes embedding batches. 

292 """ 

293 start_time = asyncio.get_event_loop().time() 

294 job_id = None # Will be extracted from context 

295 

296 # Reconstruct chunks 

297 chunk_objs = [ 

298 _EmbeddingChunk( 

299 text=c["text"], 

300 metadata=c.get("metadata"), 

301 ) 

302 for c in chunks 

303 ] 

304 

305 try: 

306 # Find job ID from progress 

307 active_jobs = self._progress_tracker.get_active_jobs() 

308 for jid in active_jobs: 

309 progress = await self._progress_tracker.get_progress(jid) 

310 if ( 

311 progress 

312 and progress.total_texts == len(chunks) 

313 and progress.status == EmbeddingStatus.PENDING 

314 ): 

315 job_id = jid 

316 break 

317 

318 if not job_id: 

319 job_id = f"batch-{datetime.now(UTC).isoformat()}" 

320 

321 # Update progress: processing 

322 await self._progress_tracker.update_progress( 

323 job_id=job_id, 

324 status=EmbeddingStatus.PROCESSING, 

325 ) 

326 

327 # Extract texts 

328 texts = [chunk.text for chunk in chunk_objs] 

329 

330 # Process in batches 

331 all_embeddings: list[list[float]] = [] 

332 texts_processed = 0 

333 cache_hits = 0 

334 cache_misses = 0 

335 

336 for i in range(0, len(texts), batch_size): 

337 batch_texts = texts[i : i + batch_size] 

338 

339 # Get embeddings (with or without cache) 

340 if use_cache and self._cache: 

341 ( 

342 batch_embeddings, 

343 hits, 

344 misses, 

345 ) = await self._cache.get_embeddings_with_cache( 

346 batch_texts, 

347 model_name, 

348 self.embedding_provider, 

349 ) 

350 cache_hits += hits 

351 cache_misses += misses 

352 else: 

353 # Generate embeddings without cache 

354 batch_embeddings = await self.embedding_provider.embed_texts( 

355 batch_texts, 

356 ) 

357 cache_misses += len(batch_texts) 

358 

359 all_embeddings.extend(batch_embeddings) 

360 texts_processed += len(batch_texts) 

361 

362 # Update progress 

363 await self._progress_tracker.update_progress( 

364 job_id=job_id, 

365 texts_processed=texts_processed, 

366 cache_hits=cache_hits, 

367 cache_misses=cache_misses, 

368 ) 

369 

370 logger.debug( 

371 "Processed embedding batch", 

372 job_id=job_id, 

373 batch_size=len(batch_texts), 

374 total_processed=texts_processed, 

375 cache_hit_rate=f"{(cache_hits / (cache_hits + cache_misses) * 100):.1f}%", 

376 ) 

377 

378 # Update progress: storing 

379 await self._progress_tracker.update_progress( 

380 job_id=job_id, 

381 status=EmbeddingStatus.STORING, 

382 ) 

383 

384 # Store embeddings in vector store 

385 await self._store_embeddings( 

386 chunks=chunk_objs, 

387 embeddings=all_embeddings, 

388 collection_name=collection_name, 

389 ) 

390 

391 # Update progress: completed 

392 await self._progress_tracker.update_progress( 

393 job_id=job_id, 

394 status=EmbeddingStatus.COMPLETED, 

395 ) 

396 

397 duration = asyncio.get_event_loop().time() - start_time 

398 

399 logger.info( 

400 "Batch embedding completed", 

401 job_id=job_id, 

402 embeddings_generated=len(all_embeddings), 

403 cache_hits=cache_hits, 

404 cache_hit_rate=f"{(cache_hits / (cache_hits + cache_misses) * 100):.1f}%", 

405 duration=f"{duration:.2f}s", 

406 ) 

407 

408 return BatchEmbeddingResult.success_result( 

409 job_id=job_id, 

410 embeddings_generated=len(all_embeddings), 

411 cache_hits=cache_hits, 

412 duration=duration, 

413 metadata={"collection": collection_name, "model": model_name}, 

414 ) 

415 

416 except Exception as e: 

417 duration = asyncio.get_event_loop().time() - start_time 

418 error_msg = str(e) 

419 

420 if job_id: 

421 await self._progress_tracker.update_progress( 

422 job_id=job_id, 

423 error=error_msg, 

424 ) 

425 

426 logger.exception( 

427 "Batch embedding failed", 

428 job_id=job_id, 

429 error=error_msg, 

430 ) 

431 

432 return BatchEmbeddingResult.failure_result( 

433 job_id=job_id or "unknown", 

434 error=error_msg, 

435 duration=duration, 

436 ) 

437 

438 async def _store_embeddings( 

439 self, 

440 chunks: list[_EmbeddingChunk], 

441 embeddings: list[list[float]], 

442 collection_name: str, 

443 ) -> None: 

444 """Store embeddings in vector store.""" 

445 # Extract texts and metadata 

446 texts = [chunk.text for chunk in chunks] 

447 metadatas = [chunk.metadata for chunk in chunks] 

448 

449 # Add to vector store with embeddings 

450 await self.vector_store.add_texts( 

451 texts=texts, 

452 embeddings=embeddings, 

453 metadatas=[m or {} for m in metadatas], 

454 collection_name=collection_name, 

455 ) 

456 

457 def get_stats(self) -> dict[str, Any]: 

458 """Get worker statistics.""" 

459 progress_stats = self._progress_tracker.get_stats() 

460 cache_size = self._cache.size() if self._cache else 0 

461 

462 return { 

463 "worker_id": self.worker_id, 

464 "running": self._running, 

465 "concurrency": self.concurrency, 

466 "active_workers": len(self._workers), 

467 "cache_size": cache_size, 

468 "cache_enabled": self.enable_cache, 

469 **progress_stats, 

470 } 

471 

472 async def health_check(self, timeout: float = 5.0) -> HealthCheckResult: 

473 """Report the health of this worker. 

474 

475 Args: 

476 timeout: Unused; present for protocol conformance. 

477 

478 Returns: 

479 HEALTHY when the worker is running, UNHEALTHY otherwise. 

480 """ 

481 status = HealthStatus.HEALTHY if self._running else HealthStatus.UNHEALTHY 

482 stats = self.get_stats() 

483 return HealthCheckResult( 

484 component=f"worker.batch_embedding.{self.worker_id}", 

485 status=status, 

486 details=stats, 

487 ) 

488 

489 async def clear_cache(self) -> None: 

490 """Clear the embedding cache.""" 

491 if self._cache: 

492 await self._cache.clear() 

493 logger.info("Cleared embedding cache", worker_id=self.worker_id)