1"""
2Embedding cache management for batch embedding operations.
3"""
4
5from __future__ import annotations
6
7import asyncio
8
9from lexigram.logging import (
10 get_logger,
11)
12
13logger = get_logger(__name__)
14
15
16class EmbeddingCache:
17 """
18 Simple in-memory cache for embeddings.
19
20 Enterprise version would use distributed cache (Redis).
21 """
22
23 def __init__(self) -> None:
24 """Initialize embedding cache."""
25 self._cache: dict[str, list[float]] = {}
26 self._lock = asyncio.Lock()
27
28 async def get(self, text: str, model_name: str) -> list[float] | None:
29 """Get cached embedding for text."""
30 cache_key = f"{model_name}:{hash(text)}"
31 async with self._lock:
32 return self._cache.get(cache_key)
33
34 async def set(self, text: str, model_name: str, embedding: list[float]) -> None:
35 """Cache embedding for text."""
36 cache_key = f"{model_name}:{hash(text)}"
37 async with self._lock:
38 self._cache[cache_key] = embedding
39
40 async def get_batch(
41 self,
42 texts: list[str],
43 model_name: str,
44 ) -> tuple[list[list[float] | None], list[tuple[int, str]]]:
45 """
46 Get cached embeddings for batch of texts.
47
48 Returns:
49 (cached_embeddings, uncached_indices_and_texts)
50 """
51 cached_embeddings: list[list[float] | None] = []
52 uncached: list[tuple[int, str]] = []
53
54 async with self._lock:
55 for i, text in enumerate(texts):
56 cache_key = f"{model_name}:{hash(text)}"
57 embedding = self._cache.get(cache_key)
58 if embedding is not None:
59 cached_embeddings.append(embedding)
60 else:
61 cached_embeddings.append(None)
62 uncached.append((i, text))
63
64 return cached_embeddings, uncached
65
66 async def set_batch(
67 self,
68 texts_and_embeddings: list[tuple[str, list[float]]],
69 model_name: str,
70 ) -> None:
71 """Cache multiple embeddings."""
72 async with self._lock:
73 for text, embedding in texts_and_embeddings:
74 cache_key = f"{model_name}:{hash(text)}"
75 self._cache[cache_key] = embedding
76
77 async def clear(self) -> None:
78 """Clear all cached embeddings."""
79 async with self._lock:
80 self._cache.clear()
81 logger.info("Cleared embedding cache")
82
83 def size(self) -> int:
84 """Get number of cached embeddings."""
85 return len(self._cache)
86
87 async def get_embeddings_with_cache(
88 self,
89 texts: list[str],
90 model_name: str,
91 embedding_provider: EmbeddingProvider, # type: ignore[name-defined]
92 ) -> tuple[list[list[float]], int, int]:
93 """
94 Get embeddings with cache lookup and generation.
95
96 Returns:
97 (embeddings, cache_hits, cache_misses)
98 """
99 embeddings: list[list[float]] = []
100 cache_hits = 0
101 cache_misses = 0
102
103 # Check cache first
104 cached_embeddings, _uncached_texts = await self.get_batch(texts, model_name)
105
106 # Fill in cached results and collect uncached texts
107 uncached_indices_and_texts: list[tuple[int, str]] = []
108 for i, (text, cached) in enumerate(zip(texts, cached_embeddings, strict=False)):
109 if cached is not None:
110 embeddings.append(cached)
111 cache_hits += 1
112 else:
113 embeddings.append([]) # Placeholder
114 uncached_indices_and_texts.append((i, text))
115 cache_misses += 1
116
117 # Generate embeddings for cache misses
118 if uncached_indices_and_texts:
119 uncached_only = [x[1] for x in uncached_indices_and_texts]
120
121 new_embeddings = await embedding_provider.embed_texts(uncached_only)
122
123 # Cache new embeddings and update results
124 await self.set_batch(
125 list(zip(uncached_only, new_embeddings, strict=False)),
126 model_name,
127 )
128
129 for (idx, _), embedding in zip(
130 uncached_indices_and_texts,
131 new_embeddings,
132 strict=False,
133 ):
134 embeddings[idx] = embedding
135
136 return embeddings, cache_hits, cache_misses