1"""Reverse HyDE generator implementation."""
2
3from __future__ import annotations
4
5from typing import Any
6
7from lexigram.ai.rag.hyde.base import AbstractHyDEGenerator
8from lexigram.ai.rag.hyde.protocols import EmbeddingClientProtocol
9from lexigram.ai.rag.hyde.types import HyDEResult, HyDEStrategy, HypotheticalDocument
10from lexigram.contracts import (
11 ChatMessage,
12 LLMClientProtocol,
13)
14
15
16class ReverseHyDEGenerator(AbstractHyDEGenerator):
17 """Reverse HyDE: Generate query from hypothetical document."""
18
19 def __init__(
20 self,
21 llm_client: LLMClientProtocol,
22 embedding_client: EmbeddingClientProtocol | None = None,
23 temperature: float = 0.5,
24 max_tokens: int = 100,
25 ):
26 """Initialize reverse HyDE generator.
27
28 Args:
29 llm_client: Client for generating queries
30 embedding_client: Optional client for generating embeddings
31 temperature: Lower temperature for precise queries
32 max_tokens: Maximum tokens per query
33 """
34 super().__init__(llm_client, embedding_client)
35 self.temperature = temperature
36 self.max_tokens = max_tokens
37
38 def _build_reverse_prompt(self, query: str) -> str:
39 """Build prompt for reverse HyDE.
40
41 Args:
42 query: Original query
43
44 Returns:
45 Prompt for generating document then query
46 """
47 return (
48 "First, write a comprehensive passage that would answer this query. "
49 "Then, generate 3-5 related queries that this passage would answer.\n\n"
50 f"Original Query: {query}\n\n"
51 "Format:\n"
52 "Passage: [your passage here]\n"
53 "Related Queries:\n"
54 "1. [query 1]\n"
55 "2. [query 2]\n"
56 "..."
57 )
58
59 async def generate(
60 self,
61 query: str,
62 num_documents: int = 1,
63 **kwargs: Any,
64 ) -> HyDEResult:
65 """Generate hypothetical document and related queries.
66
67 Args:
68 query: User query
69 num_documents: Ignored
70 **kwargs: Additional parameters
71
72 Returns:
73 HyDE result with reverse generation
74 """
75 prompt = self._build_reverse_prompt(query)
76 messages = [ChatMessage(role="user", content=prompt)]
77
78 result = await self.llm_client.complete(
79 messages,
80 temperature=self.temperature,
81 max_tokens=self.max_tokens * 3, # Need more tokens for passage + queries
82 )
83 if result.is_err():
84 raise result.unwrap_err()
85 response = result.unwrap()
86
87 content = self._extract_content(response)
88
89 # Parse passage and queries
90 passage, queries = self._parse_reverse_response(content)
91
92 doc = HypotheticalDocument(
93 content=passage,
94 query=query,
95 confidence=1.0,
96 metadata={
97 "related_queries": queries,
98 "temperature": self.temperature,
99 "max_tokens": self.max_tokens,
100 },
101 )
102
103 # Generate embedding if client available
104 aggregated_embedding = None
105 if self.embedding_client:
106 embeddings = await self._embed_documents([doc])
107 aggregated_embedding = embeddings[0] if embeddings else None
108
109 return HyDEResult(
110 query=query,
111 hypothetical_docs=[doc],
112 strategy=HyDEStrategy.REVERSE,
113 aggregated_embedding=aggregated_embedding,
114 metadata={
115 "related_queries": queries,
116 "temperature": self.temperature,
117 "max_tokens": self.max_tokens,
118 },
119 )
120
121 def _parse_reverse_response(self, content: str) -> tuple[str, list[str]]:
122 """Parse reverse HyDE response into passage and queries.
123
124 Args:
125 content: LLM response content
126
127 Returns:
128 Tuple of (passage, list of queries)
129 """
130 lines = content.strip().split("\n")
131
132 passage_lines = []
133 queries = []
134 in_passage = False
135 in_queries = False
136
137 for line in lines:
138 line_lower = line.lower().strip()
139
140 if line_lower.startswith("passage:"):
141 in_passage = True
142 in_queries = False
143 # Get passage content after "Passage:"
144 passage_content = line.split(":", 1)[1].strip()
145 if passage_content:
146 passage_lines.append(passage_content)
147 elif "related queries" in line_lower or "queries:" in line_lower:
148 in_passage = False
149 in_queries = True
150 elif in_passage:
151 passage_lines.append(line.strip())
152 elif in_queries and line.strip():
153 # Extract query (remove numbering)
154 query = line.strip()
155 # Remove leading numbers and dots
156 query = query.lstrip("0123456789.-) ").strip()
157 if query:
158 queries.append(query)
159
160 passage = " ".join(passage_lines)
161
162 return passage, queries