Coverage for src/chat_limiter/adapters.py: 84%
205 statements
« prev ^ index » next coverage.py v7.9.2, created at 2025-12-09 08:16 -0500
« prev ^ index » next coverage.py v7.9.2, created at 2025-12-09 08:16 -0500
1"""
2Provider-specific adapters for converting between our unified types and provider APIs.
3"""
5import time
6import warnings
7from abc import ABC, abstractmethod
8from typing import Any
10from .providers import Provider
11from .types import (
12 ChatCompletionRequest,
13 ChatCompletionResponse,
14 Choice,
15 Message,
16 MessageRole,
17 Usage,
18)
21class ProviderAdapter(ABC):
22 """Abstract base class for provider-specific adapters."""
24 def is_reasoning_model(self, model_name: str) -> bool:
25 """Check if the model is a reasoning model (o1, o3, o4 series)."""
26 # Handle prefixed models (e.g., "openai/o3-mini")
27 if "/" in model_name:
28 # Extract the base model name after the "/"
29 base_model = model_name.split("/", 1)[1]
30 return base_model.startswith(("o1", "o3", "o4", "gpt-5"))
32 # Handle non-prefixed models
33 return model_name.startswith(("o1", "o3", "o4", "gpt-5"))
35 @abstractmethod
36 def format_request(self, request: ChatCompletionRequest) -> dict[str, Any]:
37 """Convert our request format to provider-specific format."""
38 pass
40 @abstractmethod
41 def parse_response(
42 self,
43 response_data: dict[str, Any],
44 original_request: ChatCompletionRequest
45 ) -> ChatCompletionResponse:
46 """Convert provider response to our unified format."""
47 pass
49 @abstractmethod
50 def get_endpoint(self) -> str:
51 """Get the API endpoint for this provider."""
52 pass
55class OpenAIAdapter(ProviderAdapter):
56 """Adapter for OpenAI API."""
58 def format_request(self, request: ChatCompletionRequest) -> dict[str, Any]:
59 """Convert to OpenAI format."""
60 # Convert messages
61 messages: list[dict[str, Any]] = []
62 for msg in request.messages:
63 messages.append({
64 "role": msg.role.value,
65 "content": msg.content
66 })
68 model = request.model.strip()
69 if model.startswith("openai/"):
70 # Remove the "openai/" prefix, since we are already using the OpenAI API
71 model = model.split("openai/", 1)[1]
73 # Build request
74 openai_request: dict[str, Any] = {
75 "model": model,
76 "messages": messages,
77 }
79 # Add optional parameters
80 if request.max_tokens is not None:
81 # Use max_completion_tokens for reasoning models (o1, o3, o4)
82 if self.is_reasoning_model(model):
83 openai_request["max_completion_tokens"] = request.max_tokens
84 else:
85 openai_request["max_tokens"] = request.max_tokens
87 # Handle temperature for reasoning models
88 if self.is_reasoning_model(model):
89 # For reasoning models, default to temperature=1
90 default_temperature = 1.0
92 if request.temperature is not None:
93 # If user provided a different temperature, warn them and use temperature=1
94 if request.temperature != default_temperature:
95 warnings.warn(
96 f"WARNING: Model '{model}' is a reasoning model that requires temperature=1. "
97 f"Your specified temperature={request.temperature} will be overridden to temperature=1.",
98 UserWarning
99 )
100 print(f"WARNING: Model '{model}' is a reasoning model that requires temperature=1. "
101 f"Your specified temperature={request.temperature} will be overridden to temperature=1.")
103 # Always use temperature=1 for reasoning models
104 openai_request["temperature"] = default_temperature
105 else:
106 # For non-reasoning models, use the provided temperature
107 if request.temperature is not None:
108 openai_request["temperature"] = request.temperature
110 if request.top_p is not None:
111 openai_request["top_p"] = request.top_p
112 if request.stop is not None:
113 openai_request["stop"] = request.stop
114 if request.stream:
115 openai_request["stream"] = request.stream
116 if request.frequency_penalty is not None:
117 openai_request["frequency_penalty"] = request.frequency_penalty
118 if request.presence_penalty is not None:
119 openai_request["presence_penalty"] = request.presence_penalty
120 if request.seed is not None:
121 openai_request["seed"] = request.seed
123 # Add reasoning parameter for thinking models
124 if (request.reasoning_effort is not None and
125 self.is_reasoning_model(model)):
126 openai_request["reasoning"] = {"effort": request.reasoning_effort}
128 return openai_request
130 def parse_response(
131 self,
132 response_data: dict[str, Any],
133 original_request: ChatCompletionRequest
134 ) -> ChatCompletionResponse:
135 """Parse OpenAI response."""
136 # Check for errors first
137 success = True
138 error_message = None
140 if "error" in response_data:
141 success = False
142 error_data = response_data["error"]
143 error_message = error_data.get("message", "Unknown error")
145 choices = []
146 for choice_data in response_data.get("choices", []):
147 message_data = choice_data.get("message", {})
149 # Handle both string and content-block formats
150 raw_content = message_data.get("content", "")
151 content_text = ""
152 if isinstance(raw_content, str) and raw_content:
153 content_text = raw_content
154 elif isinstance(raw_content, list) and raw_content:
155 # Newer OpenAI responses may return a list of content blocks
156 parts: list[str] = []
157 for block in raw_content:
158 if not isinstance(block, dict):
159 continue
160 # Prefer explicit output fields used by reasoning models
161 output_text_val = block.get("output_text")
162 if isinstance(output_text_val, str) and output_text_val:
163 parts.append(output_text_val)
164 continue
165 # Fallbacks
166 text_val = block.get("text")
167 if isinstance(text_val, str) and text_val:
168 parts.append(text_val)
169 continue
170 content_val = block.get("content")
171 if isinstance(content_val, str) and content_val:
172 parts.append(content_val)
173 content_text = "".join(parts)
174 # Choice-level fallback sometimes present in reasoning responses
175 if not content_text:
176 choice_level_output = choice_data.get("output_text")
177 if isinstance(choice_level_output, str) and choice_level_output:
178 content_text = choice_level_output
180 message = Message(
181 role=MessageRole(message_data.get("role", "assistant")),
182 content=content_text
183 )
184 choice = Choice(
185 index=choice_data.get("index", 0),
186 message=message,
187 finish_reason=choice_data.get("finish_reason")
188 )
189 choices.append(choice)
191 # Parse usage
192 usage = None
193 if "usage" in response_data:
194 usage_data = response_data["usage"]
195 usage = Usage(
196 prompt_tokens=usage_data.get("prompt_tokens", 0),
197 completion_tokens=usage_data.get("completion_tokens", 0),
198 total_tokens=usage_data.get("total_tokens", 0)
199 )
201 return ChatCompletionResponse(
202 id=response_data.get("id", ""),
203 model=response_data.get("model", original_request.model),
204 choices=choices,
205 usage=usage,
206 created=response_data.get("created"),
207 success=success,
208 error_message=error_message,
209 provider="openai",
210 raw_response=response_data
211 )
213 def get_endpoint(self) -> str:
214 return "/chat/completions"
217class AnthropicAdapter(ProviderAdapter):
218 """Adapter for Anthropic API."""
220 def format_request(self, request: ChatCompletionRequest) -> dict[str, Any]:
221 """Convert to Anthropic format."""
222 # Anthropic has a different message format
223 messages: list[dict[str, Any]] = []
224 system_message: str | None = None
226 for msg in request.messages:
227 if msg.role == MessageRole.SYSTEM:
228 # Anthropic handles system messages separately
229 system_message = msg.content
230 else:
231 messages.append({
232 "role": msg.role.value,
233 "content": msg.content
234 })
236 model = request.model.strip()
237 if model.startswith("anthropic/"):
238 # Remove the "anthropic/" prefix, since we are already using the Anthropic API
239 model = model.split("anthropic/", 1)[1]
241 # Build request
242 anthropic_request: dict[str, Any] = {
243 "model": model,
244 "messages": messages,
245 "max_tokens": request.max_tokens or 1024, # Required for Anthropic
246 }
248 # Add system message if present
249 if system_message:
250 anthropic_request["system"] = system_message
252 # Add optional parameters
253 if request.temperature is not None:
254 anthropic_request["temperature"] = request.temperature
255 if request.top_p is not None:
256 anthropic_request["top_p"] = request.top_p
257 if request.stop is not None:
258 anthropic_request["stop_sequences"] = (
259 [request.stop] if isinstance(request.stop, str) else request.stop
260 )
261 if request.stream:
262 anthropic_request["stream"] = request.stream
263 if request.top_k is not None:
264 anthropic_request["top_k"] = request.top_k
265 if request.seed is not None:
266 anthropic_request["seed"] = request.seed
268 return anthropic_request
270 def parse_response(
271 self,
272 response_data: dict[str, Any],
273 original_request: ChatCompletionRequest
274 ) -> ChatCompletionResponse:
275 """Parse Anthropic response."""
276 # Check for errors first
277 success = True
278 error_message = None
280 if "error" in response_data:
281 success = False
282 error_data = response_data["error"]
283 error_message = error_data.get("message", "Unknown error")
285 # Anthropic returns content differently
286 content_blocks = response_data.get("content", [])
287 content = ""
288 if content_blocks:
289 # Extract text from content blocks
290 for block in content_blocks:
291 if block.get("type") == "text":
292 content += block.get("text", "")
294 message = Message(
295 role=MessageRole.ASSISTANT,
296 content=content
297 )
299 choice = Choice(
300 index=0,
301 message=message,
302 finish_reason=response_data.get("stop_reason")
303 )
305 # Parse usage
306 usage = None
307 if "usage" in response_data:
308 usage_data = response_data["usage"]
309 usage = Usage(
310 prompt_tokens=usage_data.get("input_tokens", 0),
311 completion_tokens=usage_data.get("output_tokens", 0),
312 total_tokens=usage_data.get("input_tokens", 0) + usage_data.get("output_tokens", 0)
313 )
315 return ChatCompletionResponse(
316 id=response_data.get("id", ""),
317 model=response_data.get("model", original_request.model),
318 choices=[choice],
319 usage=usage,
320 created=int(time.time()), # Anthropic doesn't provide created timestamp
321 success=success,
322 error_message=error_message,
323 provider="anthropic",
324 raw_response=response_data
325 )
327 def get_endpoint(self) -> str:
328 return "/messages"
331class OpenRouterAdapter(ProviderAdapter):
332 """Adapter for OpenRouter API."""
334 def format_request(self, request: ChatCompletionRequest) -> dict[str, Any]:
335 """Convert to OpenRouter format (similar to OpenAI)."""
336 # OpenRouter uses OpenAI-compatible format
337 messages: list[dict[str, Any]] = []
338 for msg in request.messages:
339 messages.append({
340 "role": msg.role.value,
341 "content": msg.content
342 })
344 model = request.model.strip()
346 # Build request
347 openrouter_request: dict[str, Any] = {
348 "model": model,
349 "messages": messages,
350 }
352 # Add optional parameters
353 if request.max_tokens is not None:
354 openrouter_request["max_tokens"] = request.max_tokens
355 if request.temperature is not None:
356 openrouter_request["temperature"] = request.temperature
357 if request.top_p is not None:
358 openrouter_request["top_p"] = request.top_p
359 if request.stop is not None:
360 openrouter_request["stop"] = request.stop
361 if request.stream:
362 openrouter_request["stream"] = request.stream
363 if request.frequency_penalty is not None:
364 openrouter_request["frequency_penalty"] = request.frequency_penalty
365 if request.presence_penalty is not None:
366 openrouter_request["presence_penalty"] = request.presence_penalty
367 if request.top_k is not None:
368 openrouter_request["top_k"] = request.top_k
369 if request.seed is not None:
370 openrouter_request["seed"] = request.seed
372 # Add reasoning parameter for thinking models
373 if (request.reasoning_effort is not None and
374 self.is_reasoning_model(model)):
375 openrouter_request["reasoning"] = {"effort": request.reasoning_effort}
377 # Add provider routing if specified
378 if request.providers is not None:
379 openrouter_request["provider"] = {
380 "order": request.providers,
381 "allow_fallbacks": False
382 }
384 return openrouter_request
386 def parse_response(
387 self,
388 response_data: dict[str, Any],
389 original_request: ChatCompletionRequest
390 ) -> ChatCompletionResponse:
391 """Parse OpenRouter response (similar to OpenAI)."""
392 # Check for errors first
393 success = True
394 error_message = None
396 if "error" in response_data:
397 success = False
398 error_data = response_data["error"]
399 error_message = error_data.get("message", "Unknown error")
401 choices = []
402 for choice_data in response_data.get("choices", []):
403 message_data = choice_data.get("message", {})
404 message = Message(
405 role=MessageRole(message_data.get("role", "assistant")),
406 content=message_data.get("content", "")
407 )
408 choice = Choice(
409 index=choice_data.get("index", 0),
410 message=message,
411 finish_reason=choice_data.get("finish_reason")
412 )
413 choices.append(choice)
415 # Parse usage
416 usage = None
417 if "usage" in response_data:
418 usage_data = response_data["usage"]
419 usage = Usage(
420 prompt_tokens=usage_data.get("prompt_tokens", 0),
421 completion_tokens=usage_data.get("completion_tokens", 0),
422 total_tokens=usage_data.get("total_tokens", 0)
423 )
425 return ChatCompletionResponse(
426 id=response_data.get("id", ""),
427 model=response_data.get("model", original_request.model),
428 choices=choices,
429 usage=usage,
430 created=response_data.get("created"),
431 success=success,
432 error_message=error_message,
433 provider="openrouter",
434 raw_response=response_data
435 )
437 def get_endpoint(self) -> str:
438 return "/chat/completions"
441# Provider adapter registry
442PROVIDER_ADAPTERS = {
443 Provider.OPENAI: OpenAIAdapter(),
444 Provider.ANTHROPIC: AnthropicAdapter(),
445 Provider.OPENROUTER: OpenRouterAdapter(),
446}
449def get_adapter(provider: Provider) -> ProviderAdapter:
450 """Get the appropriate adapter for a provider."""
451 return PROVIDER_ADAPTERS[provider]