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

59 statements  

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

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