Coverage for /home/admin/Documents/AI/applications/lexigram-dev/lexigram/experimental/ai/lexigram-ai-rag/src/lexigram/ai/rag/hyde/generators/reverse.py: 93%

54 statements  

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

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