Coverage for /home/admin/Documents/AI/applications/lexigram-dev/lexigram/experimental/ai/lexigram-ai-rag/src/lexigram/ai/rag/query/transformers.py: 95%

173 statements  

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

1from __future__ import annotations 

2 

3from collections.abc import Callable 

4 

5from lexigram.ai.rag.query.base import ( 

6 AbstractQueryTransformer, 

7 TransformationStrategy, 

8 TransformedQuery, 

9) 

10from lexigram.contracts import ( 

11 ChatMessage, 

12 LLMClientProtocol, 

13) 

14from lexigram.logging import ( 

15 get_logger, 

16) 

17 

18logger = get_logger(__name__) 

19 

20 

21class QueryExpander(AbstractQueryTransformer): 

22 """Expand queries with synonyms and related terms.""" 

23 

24 def __init__( 

25 self, 

26 expansion_terms: dict | None = None, 

27 llm_client: LLMClientProtocol | None = None, 

28 max_expansions: int = 5, 

29 include_original: bool = True, 

30 ): 

31 self.expansion_terms = expansion_terms or {} 

32 self.llm_client = llm_client 

33 self.max_expansions = max_expansions 

34 self.include_original = include_original 

35 

36 async def transform(self, query: str) -> TransformedQuery: 

37 expanded = [] 

38 if self.include_original: 

39 expanded.append(query) 

40 

41 for term, expansions in self.expansion_terms.items(): 

42 if term.lower() in query.lower(): 

43 for expansion in expansions[: self.max_expansions]: 

44 expanded_query = query.replace(term, expansion) 

45 if expanded_query not in expanded: 

46 expanded.append(expanded_query) 

47 

48 if self.llm_client and len(expanded) < self.max_expansions + 1: 

49 llm_expansions = await self._llm_expand(query) 

50 for exp in llm_expansions: 

51 if exp not in expanded and len(expanded) < self.max_expansions + 1: 

52 expanded.append(exp) 

53 

54 return TransformedQuery( 

55 original=query, 

56 transformed=expanded, 

57 strategy=TransformationStrategy.EXPANSION, 

58 metadata={"method": "hybrid" if self.llm_client else "predefined"}, 

59 ) 

60 

61 async def _llm_expand(self, query: str) -> list[str]: 

62 if not self.llm_client: 

63 return [] 

64 

65 prompt = f"""Generate {self.max_expansions} alternative phrasings of this search query. 

66Each alternative should maintain the same intent but use different words or structure. 

67 

68Query: {query} 

69 

70Alternative queries (one per line):""" 

71 

72 messages = [ChatMessage(role="user", content=prompt)] 

73 

74 try: 

75 result = await self.llm_client.complete( 

76 messages=messages, 

77 temperature=0.7, 

78 max_tokens=200, 

79 ) 

80 if result.is_err(): 

81 raise result.unwrap_err() 

82 response = result.unwrap() 

83 

84 text_response = response if isinstance(response, str) else response.content 

85 

86 lines = [ 

87 line.strip() 

88 for line in text_response.strip().split("\n") 

89 if line.strip() 

90 ] 

91 expansions = [] 

92 for line in lines: 

93 cleaned = line.lstrip("0123456789.-) ") 

94 if cleaned and cleaned != query: 

95 expansions.append(cleaned) 

96 

97 return expansions[: self.max_expansions] 

98 except Exception as e: # noqa: BLE001 — broadened intentionally; LLM expansion must not crash caller 

99 logger.error("query_expansion_failed", error=str(e), exc_info=True) 

100 return [] 

101 

102 @property 

103 def strategy(self) -> TransformationStrategy: 

104 return TransformationStrategy.EXPANSION 

105 

106 

107class MultiQueryGenerator(AbstractQueryTransformer): 

108 """Generate multiple query variations for parallel search.""" 

109 

110 def __init__( 

111 self, 

112 llm_client: LLMClientProtocol, 

113 num_queries: int = 3, 

114 include_original: bool = True, 

115 temperature: float = 0.8, 

116 ): 

117 self.llm_client = llm_client 

118 self.num_queries = num_queries 

119 self.include_original = include_original 

120 self.temperature = temperature 

121 

122 async def transform(self, query: str) -> TransformedQuery: 

123 queries = [] 

124 if self.include_original: 

125 queries.append(query) 

126 

127 prompt = f"""You are an AI assistant helping to improve search queries. 

128Generate {self.num_queries} different versions of the following query. 

129Each version should: 

130- Maintain the original intent 

131- Use different keywords and phrasings 

132- Cover different aspects of the question 

133 

134Original query: {query} 

135 

136Generate {self.num_queries} alternative queries (one per line):""" 

137 

138 messages = [ChatMessage(role="user", content=prompt)] 

139 

140 try: 

141 result = await self.llm_client.complete( 

142 messages=messages, 

143 temperature=self.temperature, 

144 max_tokens=300, 

145 ) 

146 if result.is_err(): 

147 raise result.unwrap_err() 

148 response = result.unwrap() 

149 

150 text_response = response if isinstance(response, str) else response.content 

151 

152 lines = [ 

153 line.strip() 

154 for line in text_response.strip().split("\n") 

155 if line.strip() 

156 ] 

157 for line in lines: 

158 cleaned = line.lstrip("0123456789.-) ") 

159 if cleaned and cleaned not in queries: 

160 queries.append(cleaned) 

161 if len(queries) >= self.num_queries + ( 

162 1 if self.include_original else 0 

163 ): 

164 break 

165 

166 except (ValueError, TypeError, RuntimeError, OSError): 

167 if not queries: 

168 queries.append(query) 

169 

170 return TransformedQuery( 

171 original=query, 

172 transformed=queries, 

173 strategy=TransformationStrategy.MULTI_QUERY, 

174 metadata={ 

175 "temperature": self.temperature, 

176 "target_count": self.num_queries, 

177 }, 

178 ) 

179 

180 @property 

181 def strategy(self) -> TransformationStrategy: 

182 return TransformationStrategy.MULTI_QUERY 

183 

184 

185class HyDEGenerator(AbstractQueryTransformer): 

186 """Generate hypothetical documents (HyDE) for the query.""" 

187 

188 def __init__( 

189 self, 

190 llm_client: LLMClientProtocol, 

191 num_documents: int = 1, 

192 doc_length: str = "medium", 

193 temperature: float = 0.7, 

194 ): 

195 self.llm_client = llm_client 

196 self.num_documents = num_documents 

197 self.doc_length = doc_length 

198 self.temperature = temperature 

199 

200 async def transform(self, query: str) -> TransformedQuery: 

201 documents = [] 

202 max_tokens = {"short": 150, "medium": 300, "long": 500}.get( 

203 self.doc_length, 

204 300, 

205 ) 

206 

207 prompt = f"""Write a detailed answer to the following question. 

208Provide a comprehensive response as if you were writing documentation or a textbook. 

209 

210Question: {query} 

211 

212Answer:""" 

213 

214 messages = [ChatMessage(role="user", content=prompt)] 

215 

216 for _ in range(self.num_documents): 

217 try: 

218 result = await self.llm_client.complete( 

219 messages=messages, 

220 temperature=self.temperature, 

221 max_tokens=max_tokens, 

222 ) 

223 if result.is_err(): 

224 raise result.unwrap_err() 

225 response = result.unwrap() 

226 

227 text_response = ( 

228 response if isinstance(response, str) else response.content 

229 ) 

230 

231 if text_response and text_response.strip(): 

232 documents.append(text_response.strip()) 

233 except (ValueError, TypeError, RuntimeError, OSError) as e: 

234 logger.debug("One generation failed while transforming queries: %s", e) 

235 continue 

236 

237 if not documents: 

238 documents.append(query) 

239 

240 return TransformedQuery( 

241 original=query, 

242 transformed=documents, 

243 strategy=TransformationStrategy.HYDE, 

244 metadata={ 

245 "doc_length": self.doc_length, 

246 "temperature": self.temperature, 

247 "target_count": self.num_documents, 

248 }, 

249 ) 

250 

251 @property 

252 def strategy(self) -> TransformationStrategy: 

253 return TransformationStrategy.HYDE 

254 

255 

256class QueryRewriter(AbstractQueryTransformer): 

257 """Rewrite queries for better retrieval.""" 

258 

259 def __init__( 

260 self, 

261 llm_client: LLMClientProtocol, 

262 instructions: str | None = None, 

263 temperature: float = 0.3, 

264 ): 

265 self.llm_client = llm_client 

266 self.instructions = instructions or self._default_instructions() 

267 self.temperature = temperature 

268 

269 def _default_instructions(self) -> str: 

270 return """Rewrite the query to be more clear, specific, and effective for search. 

271- Fix spelling and grammar 

272- Expand abbreviations 

273- Add relevant context 

274- Make the intent explicit 

275- Keep it concise""" 

276 

277 async def transform(self, query: str) -> TransformedQuery: 

278 prompt = f"""{self.instructions} 

279 

280Original query: {query} 

281 

282Rewritten query:""" 

283 

284 messages = [ChatMessage(role="user", content=prompt)] 

285 

286 try: 

287 result = await self.llm_client.complete( 

288 messages=messages, 

289 temperature=self.temperature, 

290 max_tokens=100, 

291 ) 

292 if result.is_err(): 

293 raise result.unwrap_err() 

294 response = result.unwrap() 

295 

296 rewritten = response if isinstance(response, str) else response.content 

297 rewritten = rewritten.strip() 

298 

299 if not rewritten or len(rewritten) < 3: 

300 rewritten = query 

301 except (ValueError, TypeError, RuntimeError, OSError): 

302 rewritten = query 

303 

304 return TransformedQuery( 

305 original=query, 

306 transformed=[rewritten], 

307 strategy=TransformationStrategy.REWRITE, 

308 metadata={"instructions_used": bool(self.instructions)}, 

309 ) 

310 

311 @property 

312 def strategy(self) -> TransformationStrategy: 

313 return TransformationStrategy.REWRITE 

314 

315 

316class CustomQueryTransformer(AbstractQueryTransformer): 

317 """Custom query transformer using user-defined function.""" 

318 

319 def __init__( 

320 self, 

321 transform_fn: Callable[[str], list[str]], 

322 strategy_name: str = "custom", 

323 ): 

324 self.transform_fn = transform_fn 

325 self.strategy_name = strategy_name 

326 

327 async def transform(self, query: str) -> TransformedQuery: 

328 try: 

329 transformed = self.transform_fn(query) 

330 if not isinstance(transformed, list): 

331 transformed = [transformed] 

332 except (ValueError, TypeError, RuntimeError, OSError): 

333 transformed = [query] 

334 

335 return TransformedQuery( 

336 original=query, 

337 transformed=transformed, 

338 strategy=TransformationStrategy.CUSTOM, 

339 metadata={"strategy_name": self.strategy_name}, 

340 ) 

341 

342 @property 

343 def strategy(self) -> TransformationStrategy: 

344 return TransformationStrategy.CUSTOM 

345 

346 

347def create_transformer( 

348 strategy: TransformationStrategy, 

349 llm_client: LLMClientProtocol | None = None, 

350 **kwargs, 

351) -> AbstractQueryTransformer: 

352 """Factory function to create query transformers.""" 

353 if strategy == TransformationStrategy.EXPANSION: 

354 return QueryExpander(llm_client=llm_client, **kwargs) 

355 if strategy == TransformationStrategy.MULTI_QUERY: 

356 if not llm_client: 

357 msg = "MULTI_QUERY requires llm_client" 

358 raise ValueError(msg) 

359 return MultiQueryGenerator(llm_client=llm_client, **kwargs) 

360 if strategy == TransformationStrategy.HYDE: 

361 if not llm_client: 

362 msg = "HYDE requires llm_client" 

363 raise ValueError(msg) 

364 return HyDEGenerator(llm_client=llm_client, **kwargs) 

365 if strategy == TransformationStrategy.REWRITE: 

366 if not llm_client: 

367 msg = "REWRITE requires llm_client" 

368 raise ValueError(msg) 

369 return QueryRewriter(llm_client=llm_client, **kwargs) 

370 msg = f"Unknown strategy: {strategy}" 

371 raise ValueError(msg)