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)