Coverage for src/chat_limiter/adapters.py: 92%
128 statements
« prev ^ index » next coverage.py v7.9.2, created at 2025-07-11 12:02 +0100
« prev ^ index » next coverage.py v7.9.2, created at 2025-07-11 12:02 +0100
1"""
2Provider-specific adapters for converting between our unified types and provider APIs.
3"""
5import time
6from abc import ABC, abstractmethod
7from typing import Any
9from .providers import Provider
10from .types import (
11 ChatCompletionRequest,
12 ChatCompletionResponse,
13 Choice,
14 Message,
15 MessageRole,
16 Usage,
17)
20class ProviderAdapter(ABC):
21 """Abstract base class for provider-specific adapters."""
23 @abstractmethod
24 def format_request(self, request: ChatCompletionRequest) -> dict[str, Any]:
25 """Convert our request format to provider-specific format."""
26 pass
28 @abstractmethod
29 def parse_response(
30 self,
31 response_data: dict[str, Any],
32 original_request: ChatCompletionRequest
33 ) -> ChatCompletionResponse:
34 """Convert provider response to our unified format."""
35 pass
37 @abstractmethod
38 def get_endpoint(self) -> str:
39 """Get the API endpoint for this provider."""
40 pass
43class OpenAIAdapter(ProviderAdapter):
44 """Adapter for OpenAI API."""
46 def format_request(self, request: ChatCompletionRequest) -> dict[str, Any]:
47 """Convert to OpenAI format."""
48 # Convert messages
49 messages: list[dict[str, Any]] = []
50 for msg in request.messages:
51 messages.append({
52 "role": msg.role.value,
53 "content": msg.content
54 })
56 # Build request
57 openai_request: dict[str, Any] = {
58 "model": request.model,
59 "messages": messages,
60 }
62 # Add optional parameters
63 if request.max_tokens is not None:
64 openai_request["max_tokens"] = request.max_tokens
65 if request.temperature is not None:
66 openai_request["temperature"] = request.temperature
67 if request.top_p is not None:
68 openai_request["top_p"] = request.top_p
69 if request.stop is not None:
70 openai_request["stop"] = request.stop
71 if request.stream:
72 openai_request["stream"] = request.stream
73 if request.frequency_penalty is not None:
74 openai_request["frequency_penalty"] = request.frequency_penalty
75 if request.presence_penalty is not None:
76 openai_request["presence_penalty"] = request.presence_penalty
78 return openai_request
80 def parse_response(
81 self,
82 response_data: dict[str, Any],
83 original_request: ChatCompletionRequest
84 ) -> ChatCompletionResponse:
85 """Parse OpenAI response."""
86 choices = []
87 for choice_data in response_data.get("choices", []):
88 message_data = choice_data.get("message", {})
89 message = Message(
90 role=MessageRole(message_data.get("role", "assistant")),
91 content=message_data.get("content", "")
92 )
93 choice = Choice(
94 index=choice_data.get("index", 0),
95 message=message,
96 finish_reason=choice_data.get("finish_reason")
97 )
98 choices.append(choice)
100 # Parse usage
101 usage = None
102 if "usage" in response_data:
103 usage_data = response_data["usage"]
104 usage = Usage(
105 prompt_tokens=usage_data.get("prompt_tokens", 0),
106 completion_tokens=usage_data.get("completion_tokens", 0),
107 total_tokens=usage_data.get("total_tokens", 0)
108 )
110 return ChatCompletionResponse(
111 id=response_data.get("id", ""),
112 model=response_data.get("model", original_request.model),
113 choices=choices,
114 usage=usage,
115 created=response_data.get("created"),
116 provider="openai",
117 raw_response=response_data
118 )
120 def get_endpoint(self) -> str:
121 return "/chat/completions"
124class AnthropicAdapter(ProviderAdapter):
125 """Adapter for Anthropic API."""
127 def format_request(self, request: ChatCompletionRequest) -> dict[str, Any]:
128 """Convert to Anthropic format."""
129 # Anthropic has a different message format
130 messages: list[dict[str, Any]] = []
131 system_message: str | None = None
133 for msg in request.messages:
134 if msg.role == MessageRole.SYSTEM:
135 # Anthropic handles system messages separately
136 system_message = msg.content
137 else:
138 messages.append({
139 "role": msg.role.value,
140 "content": msg.content
141 })
143 # Build request
144 anthropic_request: dict[str, Any] = {
145 "model": request.model,
146 "messages": messages,
147 "max_tokens": request.max_tokens or 1024, # Required for Anthropic
148 }
150 # Add system message if present
151 if system_message:
152 anthropic_request["system"] = system_message
154 # Add optional parameters
155 if request.temperature is not None:
156 anthropic_request["temperature"] = request.temperature
157 if request.top_p is not None:
158 anthropic_request["top_p"] = request.top_p
159 if request.stop is not None:
160 anthropic_request["stop_sequences"] = (
161 [request.stop] if isinstance(request.stop, str) else request.stop
162 )
163 if request.stream:
164 anthropic_request["stream"] = request.stream
165 if request.top_k is not None:
166 anthropic_request["top_k"] = request.top_k
168 return anthropic_request
170 def parse_response(
171 self,
172 response_data: dict[str, Any],
173 original_request: ChatCompletionRequest
174 ) -> ChatCompletionResponse:
175 """Parse Anthropic response."""
176 # Anthropic returns content differently
177 content_blocks = response_data.get("content", [])
178 content = ""
179 if content_blocks:
180 # Extract text from content blocks
181 for block in content_blocks:
182 if block.get("type") == "text":
183 content += block.get("text", "")
185 message = Message(
186 role=MessageRole.ASSISTANT,
187 content=content
188 )
190 choice = Choice(
191 index=0,
192 message=message,
193 finish_reason=response_data.get("stop_reason")
194 )
196 # Parse usage
197 usage = None
198 if "usage" in response_data:
199 usage_data = response_data["usage"]
200 usage = Usage(
201 prompt_tokens=usage_data.get("input_tokens", 0),
202 completion_tokens=usage_data.get("output_tokens", 0),
203 total_tokens=usage_data.get("input_tokens", 0) + usage_data.get("output_tokens", 0)
204 )
206 return ChatCompletionResponse(
207 id=response_data.get("id", ""),
208 model=response_data.get("model", original_request.model),
209 choices=[choice],
210 usage=usage,
211 created=int(time.time()), # Anthropic doesn't provide created timestamp
212 provider="anthropic",
213 raw_response=response_data
214 )
216 def get_endpoint(self) -> str:
217 return "/messages"
220class OpenRouterAdapter(ProviderAdapter):
221 """Adapter for OpenRouter API."""
223 def format_request(self, request: ChatCompletionRequest) -> dict[str, Any]:
224 """Convert to OpenRouter format (similar to OpenAI)."""
225 # OpenRouter uses OpenAI-compatible format
226 messages: list[dict[str, Any]] = []
227 for msg in request.messages:
228 messages.append({
229 "role": msg.role.value,
230 "content": msg.content
231 })
233 # Build request
234 openrouter_request: dict[str, Any] = {
235 "model": request.model,
236 "messages": messages,
237 }
239 # Add optional parameters
240 if request.max_tokens is not None:
241 openrouter_request["max_tokens"] = request.max_tokens
242 if request.temperature is not None:
243 openrouter_request["temperature"] = request.temperature
244 if request.top_p is not None:
245 openrouter_request["top_p"] = request.top_p
246 if request.stop is not None:
247 openrouter_request["stop"] = request.stop
248 if request.stream:
249 openrouter_request["stream"] = request.stream
250 if request.frequency_penalty is not None:
251 openrouter_request["frequency_penalty"] = request.frequency_penalty
252 if request.presence_penalty is not None:
253 openrouter_request["presence_penalty"] = request.presence_penalty
254 if request.top_k is not None:
255 openrouter_request["top_k"] = request.top_k
257 return openrouter_request
259 def parse_response(
260 self,
261 response_data: dict[str, Any],
262 original_request: ChatCompletionRequest
263 ) -> ChatCompletionResponse:
264 """Parse OpenRouter response (similar to OpenAI)."""
265 choices = []
266 for choice_data in response_data.get("choices", []):
267 message_data = choice_data.get("message", {})
268 message = Message(
269 role=MessageRole(message_data.get("role", "assistant")),
270 content=message_data.get("content", "")
271 )
272 choice = Choice(
273 index=choice_data.get("index", 0),
274 message=message,
275 finish_reason=choice_data.get("finish_reason")
276 )
277 choices.append(choice)
279 # Parse usage
280 usage = None
281 if "usage" in response_data:
282 usage_data = response_data["usage"]
283 usage = Usage(
284 prompt_tokens=usage_data.get("prompt_tokens", 0),
285 completion_tokens=usage_data.get("completion_tokens", 0),
286 total_tokens=usage_data.get("total_tokens", 0)
287 )
289 return ChatCompletionResponse(
290 id=response_data.get("id", ""),
291 model=response_data.get("model", original_request.model),
292 choices=choices,
293 usage=usage,
294 created=response_data.get("created"),
295 provider="openrouter",
296 raw_response=response_data
297 )
299 def get_endpoint(self) -> str:
300 return "/chat/completions"
303# Provider adapter registry
304PROVIDER_ADAPTERS = {
305 Provider.OPENAI: OpenAIAdapter(),
306 Provider.ANTHROPIC: AnthropicAdapter(),
307 Provider.OPENROUTER: OpenRouterAdapter(),
308}
311def get_adapter(provider: Provider) -> ProviderAdapter:
312 """Get the appropriate adapter for a provider."""
313 return PROVIDER_ADAPTERS[provider]