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
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-25 07:19 +0800
1from __future__ import annotations
3from collections.abc import Callable
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)
18logger = get_logger(__name__)
21class QueryExpander(AbstractQueryTransformer):
22 """Expand queries with synonyms and related terms."""
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
36 async def transform(self, query: str) -> TransformedQuery:
37 expanded = []
38 if self.include_original:
39 expanded.append(query)
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)
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)
54 return TransformedQuery(
55 original=query,
56 transformed=expanded,
57 strategy=TransformationStrategy.EXPANSION,
58 metadata={"method": "hybrid" if self.llm_client else "predefined"},
59 )
61 async def _llm_expand(self, query: str) -> list[str]:
62 if not self.llm_client:
63 return []
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.
68Query: {query}
70Alternative queries (one per line):"""
72 messages = [ChatMessage(role="user", content=prompt)]
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()
84 text_response = response if isinstance(response, str) else response.content
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)
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 []
102 @property
103 def strategy(self) -> TransformationStrategy:
104 return TransformationStrategy.EXPANSION
107class MultiQueryGenerator(AbstractQueryTransformer):
108 """Generate multiple query variations for parallel search."""
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
122 async def transform(self, query: str) -> TransformedQuery:
123 queries = []
124 if self.include_original:
125 queries.append(query)
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
134Original query: {query}
136Generate {self.num_queries} alternative queries (one per line):"""
138 messages = [ChatMessage(role="user", content=prompt)]
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()
150 text_response = response if isinstance(response, str) else response.content
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
166 except (ValueError, TypeError, RuntimeError, OSError):
167 if not queries:
168 queries.append(query)
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 )
180 @property
181 def strategy(self) -> TransformationStrategy:
182 return TransformationStrategy.MULTI_QUERY
185class HyDEGenerator(AbstractQueryTransformer):
186 """Generate hypothetical documents (HyDE) for the query."""
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
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 )
207 prompt = f"""Write a detailed answer to the following question.
208Provide a comprehensive response as if you were writing documentation or a textbook.
210Question: {query}
212Answer:"""
214 messages = [ChatMessage(role="user", content=prompt)]
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()
227 text_response = (
228 response if isinstance(response, str) else response.content
229 )
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
237 if not documents:
238 documents.append(query)
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 )
251 @property
252 def strategy(self) -> TransformationStrategy:
253 return TransformationStrategy.HYDE
256class QueryRewriter(AbstractQueryTransformer):
257 """Rewrite queries for better retrieval."""
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
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"""
277 async def transform(self, query: str) -> TransformedQuery:
278 prompt = f"""{self.instructions}
280Original query: {query}
282Rewritten query:"""
284 messages = [ChatMessage(role="user", content=prompt)]
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()
296 rewritten = response if isinstance(response, str) else response.content
297 rewritten = rewritten.strip()
299 if not rewritten or len(rewritten) < 3:
300 rewritten = query
301 except (ValueError, TypeError, RuntimeError, OSError):
302 rewritten = query
304 return TransformedQuery(
305 original=query,
306 transformed=[rewritten],
307 strategy=TransformationStrategy.REWRITE,
308 metadata={"instructions_used": bool(self.instructions)},
309 )
311 @property
312 def strategy(self) -> TransformationStrategy:
313 return TransformationStrategy.REWRITE
316class CustomQueryTransformer(AbstractQueryTransformer):
317 """Custom query transformer using user-defined function."""
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
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]
335 return TransformedQuery(
336 original=query,
337 transformed=transformed,
338 strategy=TransformationStrategy.CUSTOM,
339 metadata={"strategy_name": self.strategy_name},
340 )
342 @property
343 def strategy(self) -> TransformationStrategy:
344 return TransformationStrategy.CUSTOM
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)