1from __future__ import annotations
2
3from datetime import UTC, datetime
4from typing import Any
5
6from lexigram.ai.rag.reasoning.base import (
7 AbstractReasoner,
8 ReasoningResult,
9 ReasoningStep,
10 ReasoningStrategy,
11)
12from lexigram.contracts import (
13 ChatMessage,
14 LLMClientProtocol,
15)
16from lexigram.contracts.data.vector.protocols import VectorCollectionProtocol
17
18
19class QueryDecomposer(AbstractReasoner):
20 """Decompose complex queries into simpler sub-queries.
21
22 This reasoner breaks down complex questions into simpler sub-questions
23 that can be answered independently, then combines the results.
24 """
25
26 def __init__(
27 self,
28 llm_client: LLMClientProtocol,
29 vector_store: VectorCollectionProtocol,
30 max_sub_queries: int = 5,
31 top_k_per_query: int = 2,
32 temperature: float = 0.3,
33 ):
34 """Initialize query decomposer."""
35 self.llm_client = llm_client
36 self.vector_store = vector_store
37 self.max_sub_queries = max_sub_queries
38 self.top_k_per_query = top_k_per_query
39 self.temperature = temperature
40
41 async def reason(
42 self,
43 query: str,
44 initial_context: list[Any] | None = None,
45 **kwargs,
46 ) -> ReasoningResult:
47 """Perform query decomposition reasoning."""
48 # Step 1: Decompose query into sub-queries
49 sub_queries = await self._decompose_query(query)
50
51 steps: list[ReasoningStep] = []
52
53 # Step 2: Answer each sub-query
54 for i, sub_query in enumerate(sub_queries, 1):
55 # Retrieve context for sub-query
56 retrieved_docs = await self.vector_store.search( # type: ignore[call-arg]
57 query=sub_query, # type: ignore[arg-type]
58 limit=self.top_k_per_query,
59 )
60
61 # Extract text
62 context_texts: list[str] = []
63 for doc in retrieved_docs:
64 if hasattr(doc, "content"):
65 text = doc.content
66 if text is not None:
67 context_texts.append(text)
68 elif isinstance(doc, dict) and "content" in doc:
69 context_texts.append(doc["content"])
70 elif isinstance(doc, str):
71 context_texts.append(doc)
72
73 # Answer sub-query
74 answer = await self._answer_sub_query(sub_query, context_texts)
75
76 step = ReasoningStep(
77 step_number=i,
78 question=sub_query,
79 context=retrieved_docs,
80 reasoning=f"Sub-query {i} of {len(sub_queries)}",
81 answer=answer,
82 confidence=0.7,
83 metadata={"num_docs": len(retrieved_docs)},
84 )
85 steps.append(step)
86
87 # Step 3: Synthesize final answer from sub-answers
88 final_answer = await self._synthesize_answer(query, steps)
89
90 return ReasoningResult(
91 query=query,
92 final_answer=final_answer,
93 steps=steps,
94 strategy=ReasoningStrategy.DECOMPOSITION,
95 total_hops=len(steps),
96 overall_confidence=0.75,
97 metadata={
98 "max_sub_queries": self.max_sub_queries,
99 "timestamp": datetime.now(UTC).isoformat(),
100 },
101 )
102
103 async def _decompose_query(self, query: str) -> list[str]:
104 """Decompose complex query into sub-queries."""
105 prompt = f"""Break down this complex question into simpler sub-questions that can be answered independently:
106
107Question: {query}
108
109Provide up to {self.max_sub_queries} sub-questions, one per line, numbered:
1101. <first sub-question>
1112. <second sub-question>
112...
113"""
114
115 result = await self.llm_client.complete(
116 messages=[
117 ChatMessage(
118 role="system",
119 content="You are a helpful assistant that breaks down complex questions into simpler sub-questions.",
120 ),
121 ChatMessage(role="user", content=prompt),
122 ],
123 temperature=self.temperature,
124 max_tokens=300,
125 )
126 if result.is_err():
127 raise result.unwrap_err()
128 response = result.unwrap()
129
130 # Extract text
131 if hasattr(response, "content"):
132 text = response.content
133 elif hasattr(response, "choices") and response.choices:
134 text = response.choices[0].message.content
135 elif isinstance(response, dict) and "content" in response:
136 text = response["content"]
137 else:
138 text = str(response)
139
140 # Parse numbered lines
141 sub_queries = []
142 for raw_line in text.strip().split("\n"):
143 line = raw_line.strip()
144 # Remove numbering
145 if line and (line[0].isdigit() or line.startswith(("-", "•"))):
146 # Remove number and dot/dash
147 parts = line.split(".", 1) if "." in line else line.split(")", 1)
148 sub_query = parts[1].strip() if len(parts) > 1 else line[1:].strip()
149
150 if sub_query:
151 sub_queries.append(sub_query)
152
153 return sub_queries[: self.max_sub_queries]
154
155 async def _answer_sub_query(self, sub_query: str, context: list[str]) -> str:
156 """Answer a single sub-query."""
157 context_str = "\n\n".join(f"[{i + 1}] {ctx}" for i, ctx in enumerate(context))
158
159 prompt = f"""Context:
160{context_str}
161
162Question: {sub_query}
163
164Provide a concise answer based on the context:"""
165
166 result = await self.llm_client.complete(
167 messages=[
168 ChatMessage(
169 role="system",
170 content="You are a helpful assistant that answers questions based on provided context.",
171 ),
172 ChatMessage(role="user", content=prompt),
173 ],
174 temperature=self.temperature,
175 max_tokens=200,
176 )
177 if result.is_err():
178 raise result.unwrap_err()
179 response = result.unwrap()
180
181 # Extract text
182 if hasattr(response, "content"):
183 return response.content
184 if hasattr(response, "choices") and response.choices:
185 return response.choices[0].message.content
186 if isinstance(response, dict) and "content" in response:
187 return response["content"]
188 return str(response)
189
190 async def _synthesize_answer(self, query: str, steps: list[ReasoningStep]) -> str:
191 """Synthesize final answer from sub-query answers."""
192 prompt = f"Original Question: {query}\n\n"
193 prompt += "Sub-question Answers:\n"
194 for step in steps:
195 prompt += f"{step.step_number}. {step.question}\n"
196 prompt += f" Answer: {step.answer}\n\n"
197
198 prompt += "Based on these sub-question answers, provide a comprehensive answer to the original question:"
199
200 result = await self.llm_client.complete(
201 messages=[
202 ChatMessage(
203 role="system",
204 content="You are a helpful assistant that synthesizes information from multiple sources.",
205 ),
206 ChatMessage(role="user", content=prompt),
207 ],
208 temperature=self.temperature,
209 max_tokens=400,
210 )
211 if result.is_err():
212 raise result.unwrap_err()
213 response = result.unwrap()
214
215 # Extract text
216 if hasattr(response, "content"):
217 return response.content
218 if hasattr(response, "choices") and response.choices:
219 return response.choices[0].message.content
220 if isinstance(response, dict) and "content" in response:
221 return response["content"]
222 return str(response)