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

51 statements  

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

1"""HyDE strategy registry for query generation strategies.""" 

2 

3from __future__ import annotations 

4 

5from typing import Any, Protocol 

6 

7from lexigram.ai.rag.hyde.generators import ( 

8 MultipleHyDEGenerator, 

9 ReverseHyDEGenerator, 

10 SingleHyDEGenerator, 

11 WeightedHyDEGenerator, 

12) 

13from lexigram.ai.rag.hyde.types import HyDEStrategy 

14from lexigram.contracts.ai import EmbeddingClientProtocol, LLMClientProtocol 

15 

16 

17class HyDEStrategyHandler(Protocol): 

18 """Protocol for HyDE strategy handlers.""" 

19 

20 def can_handle(self, strategy: HyDEStrategy) -> bool: 

21 """Check if this handler can handle the strategy.""" 

22 ... 

23 

24 async def create_and_generate( 

25 self, 

26 strategy: HyDEStrategy, 

27 llm_client: LLMClientProtocol, 

28 embedding_client: EmbeddingClientProtocol, 

29 query: str, 

30 num_documents: int, 

31 kwargs: dict[str, Any], 

32 ) -> Any: 

33 """Create generator and generate queries.""" 

34 ... 

35 

36 

37class SingleHyDEStrategyHandler: 

38 """Handler for SINGLE HyDE strategy.""" 

39 

40 def can_handle(self, strategy: HyDEStrategy) -> bool: 

41 return strategy == HyDEStrategy.SINGLE 

42 

43 async def create_and_generate( 

44 self, 

45 strategy: HyDEStrategy, 

46 llm_client: LLMClientProtocol, 

47 embedding_client: EmbeddingClientProtocol, 

48 query: str, 

49 num_documents: int, 

50 kwargs: dict[str, Any], 

51 ) -> Any: 

52 generator = SingleHyDEGenerator(llm_client, embedding_client) 

53 return await generator.generate(query, num_documents=num_documents, **kwargs) 

54 

55 

56class MultipleHyDEStrategyHandler: 

57 """Handler for MULTIPLE HyDE strategy.""" 

58 

59 def can_handle(self, strategy: HyDEStrategy) -> bool: 

60 return strategy == HyDEStrategy.MULTIPLE 

61 

62 async def create_and_generate( 

63 self, 

64 strategy: HyDEStrategy, 

65 llm_client: LLMClientProtocol, 

66 embedding_client: EmbeddingClientProtocol, 

67 query: str, 

68 num_documents: int, 

69 kwargs: dict[str, Any], 

70 ) -> Any: 

71 generator = MultipleHyDEGenerator(llm_client, embedding_client) 

72 return await generator.generate(query, num_documents=num_documents, **kwargs) 

73 

74 

75class WeightedHyDEStrategyHandler: 

76 """Handler for WEIGHTED HyDE strategy.""" 

77 

78 def can_handle(self, strategy: HyDEStrategy) -> bool: 

79 return strategy == HyDEStrategy.WEIGHTED 

80 

81 async def create_and_generate( 

82 self, 

83 strategy: HyDEStrategy, 

84 llm_client: LLMClientProtocol, 

85 embedding_client: EmbeddingClientProtocol, 

86 query: str, 

87 num_documents: int, 

88 kwargs: dict[str, Any], 

89 ) -> Any: 

90 if not embedding_client: 

91 msg = "Weighted HyDE requires embedding_client" 

92 raise ValueError(msg) 

93 generator = WeightedHyDEGenerator(llm_client, embedding_client) 

94 return await generator.generate(query, num_documents=num_documents, **kwargs) 

95 

96 

97class ReverseHyDEStrategyHandler: 

98 """Handler for REVERSE HyDE strategy.""" 

99 

100 def can_handle(self, strategy: HyDEStrategy) -> bool: 

101 return strategy == HyDEStrategy.REVERSE 

102 

103 async def create_and_generate( 

104 self, 

105 strategy: HyDEStrategy, 

106 llm_client: LLMClientProtocol, 

107 embedding_client: EmbeddingClientProtocol, 

108 query: str, 

109 num_documents: int, 

110 kwargs: dict[str, Any], 

111 ) -> Any: 

112 generator = ReverseHyDEGenerator(llm_client, embedding_client) 

113 return await generator.generate(query, num_documents=num_documents, **kwargs) 

114 

115 

116class HyDEStrategyRegistry: 

117 """Central registry for HyDE strategy handlers.""" 

118 

119 def __init__(self) -> None: 

120 self._handlers: list[HyDEStrategyHandler] = [] 

121 

122 @classmethod 

123 def with_defaults(cls) -> HyDEStrategyRegistry: 

124 """Create a registry pre-populated with all built-in strategy handlers.""" 

125 registry = cls() 

126 registry._handlers = [ 

127 SingleHyDEStrategyHandler(), 

128 MultipleHyDEStrategyHandler(), 

129 WeightedHyDEStrategyHandler(), 

130 ReverseHyDEStrategyHandler(), 

131 ] 

132 return registry 

133 

134 def register(self, handler: HyDEStrategyHandler) -> None: 

135 """Register a new strategy handler.""" 

136 self._handlers.insert(0, handler) 

137 

138 async def generate( 

139 self, 

140 strategy: HyDEStrategy, 

141 llm_client: LLMClientProtocol, 

142 embedding_client: EmbeddingClientProtocol, 

143 query: str, 

144 num_documents: int, 

145 kwargs: dict[str, Any], 

146 ) -> Any: 

147 """Generate queries using the appropriate strategy.""" 

148 for handler in self._handlers: 

149 if handler.can_handle(strategy): 

150 return await handler.create_and_generate( 

151 strategy, 

152 llm_client, 

153 embedding_client, 

154 query, 

155 num_documents, 

156 kwargs, 

157 ) 

158 msg = f"Unknown HyDE strategy: {strategy}" 

159 raise ValueError(msg)