Coverage for /home/admin/Documents/AI/applications/lexigram-dev/lexigram/experimental/ai/lexigram-ai-rag/src/lexigram/ai/rag/cache/manager.py: 26%

102 statements  

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

1from __future__ import annotations 

2 

3from typing import TYPE_CHECKING, Any, cast 

4 

5if TYPE_CHECKING: 

6 from lexigram.contracts import CacheBackendProtocol 

7 

8from lexigram.ai.rag.cache.base import RAGCacheConfig 

9from lexigram.ai.rag.cache.keys import CacheKeyBuilder 

10from lexigram.ai.rag.cache.stats import RAGCacheStats 

11from lexigram.logging import ( 

12 get_logger, 

13) 

14 

15logger = get_logger(__name__) 

16 

17 

18class RAGCache: 

19 """Caching layer for RAG operations delegating to platform CacheBackendProtocol.""" 

20 

21 def __init__( 

22 self, backend: CacheBackendProtocol, config: RAGCacheConfig | None = None 

23 ): 

24 """Initialize RAG cache. 

25 

26 Args: 

27 backend: The platform's cache backend 

28 config: Optional configuration 

29 """ 

30 self.backend = backend 

31 self.config = config or RAGCacheConfig() 

32 self._stats = RAGCacheStats() 

33 

34 async def cache_embedding( 

35 self, 

36 text: str, 

37 model: str, 

38 embedding: list[float], 

39 ) -> None: 

40 """Cache an embedding vector.""" 

41 key = CacheKeyBuilder.build_embedding_key(text, model, self.config.key_prefix) 

42 result = await self.backend.set( 

43 key, 

44 embedding, 

45 ttl=self.config.embedding_ttl, 

46 ) 

47 if result: 

48 self._stats.sets += 1 

49 

50 async def get_embedding( 

51 self, 

52 text: str, 

53 model: str, 

54 ) -> list[float] | None: 

55 """Get cached embedding vector.""" 

56 key = CacheKeyBuilder.build_embedding_key(text, model, self.config.key_prefix) 

57 value = cast("Any | None", await self.backend.get(key)) 

58 

59 if value is None: 

60 self._stats.misses += 1 

61 return None 

62 

63 self._stats.hits += 1 

64 return value 

65 

66 async def cache_retrieval( 

67 self, 

68 query: str, 

69 results: list[dict[str, Any]], 

70 params: dict[str, Any] | None = None, 

71 ) -> None: 

72 """Cache retrieval results.""" 

73 key = CacheKeyBuilder.build_retrieval_key(query, params, self.config.key_prefix) 

74 result = await self.backend.set( 

75 key, 

76 results, 

77 ttl=self.config.retrieval_ttl, 

78 ) 

79 if result: 

80 self._stats.sets += 1 

81 

82 async def get_retrieval( 

83 self, 

84 query: str, 

85 params: dict[str, Any] | None = None, 

86 ) -> list[dict[str, Any]] | None: 

87 """Get cached retrieval results.""" 

88 key = CacheKeyBuilder.build_retrieval_key(query, params, self.config.key_prefix) 

89 value = cast("Any | None", await self.backend.get(key)) 

90 

91 if value is None: 

92 self._stats.errors += 1 

93 return None 

94 

95 self._stats.hits += 1 

96 return value 

97 

98 async def cache_document( 

99 self, 

100 doc_id: str, 

101 document: dict[str, Any], 

102 config: dict[str, Any] | None = None, 

103 ) -> None: 

104 """Cache a preprocessed document.""" 

105 key = CacheKeyBuilder.build_document_key(doc_id, config, self.config.key_prefix) 

106 result = await self.backend.set( 

107 key, 

108 document, 

109 ttl=self.config.document_ttl, 

110 ) 

111 if result: 

112 self._stats.sets += 1 

113 

114 async def get_document( 

115 self, 

116 doc_id: str, 

117 config: dict[str, Any] | None = None, 

118 ) -> dict[str, Any] | None: 

119 """Get cached preprocessed document.""" 

120 key = CacheKeyBuilder.build_document_key(doc_id, config, self.config.key_prefix) 

121 value = cast("Any | None", await self.backend.get(key)) 

122 

123 if value is None: 

124 self._stats.errors += 1 

125 return None 

126 

127 self._stats.hits += 1 

128 return value 

129 

130 async def cache_reranking( 

131 self, 

132 query: str, 

133 document_ids: list[str], 

134 model: str, 

135 scores: list[float], 

136 ) -> None: 

137 """Cache reranking results.""" 

138 key = CacheKeyBuilder.build_reranking_key( 

139 query, 

140 document_ids, 

141 model, 

142 self.config.key_prefix, 

143 ) 

144 result = await self.backend.set( 

145 key, 

146 scores, 

147 ttl=self.config.reranking_ttl, 

148 ) 

149 if result: 

150 self._stats.sets += 1 

151 

152 async def get_reranking( 

153 self, 

154 query: str, 

155 document_ids: list[str], 

156 model: str, 

157 ) -> list[float] | None: 

158 """Get cached reranking scores.""" 

159 key = CacheKeyBuilder.build_reranking_key( 

160 query, 

161 document_ids, 

162 model, 

163 self.config.key_prefix, 

164 ) 

165 value = cast("Any | None", await self.backend.get(key)) 

166 

167 if value is None: 

168 self._stats.errors += 1 

169 return None 

170 

171 self._stats.hits += 1 

172 return value 

173 

174 async def cache_query_transformation( 

175 self, 

176 query: str, 

177 transformation_type: str, 

178 transformed: str | list[str], 

179 params: dict[str, Any] | None = None, 

180 ) -> None: 

181 """Cache query transformation results.""" 

182 key = CacheKeyBuilder.build_query_transformation_key( 

183 query, 

184 transformation_type, 

185 params, 

186 self.config.key_prefix, 

187 ) 

188 result = await self.backend.set( 

189 key, 

190 transformed, 

191 ttl=self.config.query_transformation_ttl, 

192 ) 

193 if result: 

194 self._stats.sets += 1 

195 

196 async def get_query_transformation( 

197 self, 

198 query: str, 

199 transformation_type: str, 

200 params: dict[str, Any] | None = None, 

201 ) -> str | list[str] | None: 

202 """Get cached query transformation.""" 

203 key = CacheKeyBuilder.build_query_transformation_key( 

204 query, 

205 transformation_type, 

206 params, 

207 self.config.key_prefix, 

208 ) 

209 value = cast("Any | None", await self.backend.get(key)) 

210 

211 if value is None: 

212 self._stats.errors += 1 

213 return None 

214 

215 self._stats.hits += 1 

216 return value 

217 

218 async def invalidate(self, key: str) -> bool: 

219 """Invalidate a specific cache entry.""" 

220 result = await self.backend.delete(key) 

221 if result: 

222 self._stats.deletes += 1 

223 return True 

224 return False 

225 

226 async def invalidate_pattern(self, pattern: str) -> int: 

227 """Invalidate cache entries matching a pattern.""" 

228 if hasattr(self.backend, "invalidate_pattern"): 

229 result = await self.backend.invalidate_pattern(pattern) 

230 if result: 

231 count = result 

232 self._stats.deletes += count 

233 return count 

234 

235 logger.warning( 

236 "Pattern-based invalidation not supported by CacheBackendProtocol" 

237 ) 

238 return 0 

239 

240 async def clear(self) -> None: 

241 """Clear all cache entries.""" 

242 await self.backend.clear() 

243 self._stats.deletes += 1 

244 

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

246 """Get cache statistics.""" 

247 total_entries = 0 

248 if hasattr(self.backend, "_data"): 

249 total_entries = len(self.backend._data) 

250 

251 return { 

252 "hits": self._stats.hits, 

253 "misses": self._stats.misses, 

254 "sets": self._stats.sets, 

255 "deletes": self._stats.deletes, 

256 "errors": self._stats.errors, 

257 "hit_rate": self._stats.hit_rate, 

258 "total_operations": self._stats.total_operations, 

259 "total_entries": total_entries, 

260 } 

261 

262 async def cleanup_expired(self) -> int: 

263 """Remove expired entries from cache (handled by backend).""" 

264 return 0