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)