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

1""" 

2Provider-specific adapters for converting between our unified types and provider APIs. 

3""" 

4 

5import time 

6import warnings 

7from abc import ABC, abstractmethod 

8from typing import Any 

9 

10from .providers import Provider 

11from .types import ( 

12 ChatCompletionRequest, 

13 ChatCompletionResponse, 

14 Choice, 

15 Message, 

16 MessageRole, 

17 Usage, 

18) 

19 

20 

21class ProviderAdapter(ABC): 

22 """Abstract base class for provider-specific adapters.""" 

23 

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")) 

31 

32 # Handle non-prefixed models 

33 return model_name.startswith(("o1", "o3", "o4", "gpt-5")) 

34 

35 @abstractmethod 

36 def format_request(self, request: ChatCompletionRequest) -> dict[str, Any]: 

37 """Convert our request format to provider-specific format.""" 

38 pass 

39 

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 

48 

49 @abstractmethod 

50 def get_endpoint(self) -> str: 

51 """Get the API endpoint for this provider.""" 

52 pass 

53 

54 

55class OpenAIAdapter(ProviderAdapter): 

56 """Adapter for OpenAI API.""" 

57 

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 }) 

67 

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] 

72 

73 # Build request 

74 openai_request: dict[str, Any] = { 

75 "model": model, 

76 "messages": messages, 

77 } 

78 

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 

86 

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 

91 

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.") 

102 

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 

109 

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 

122 

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} 

127 

128 return openai_request 

129 

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 

139 

140 if "error" in response_data: 

141 success = False 

142 error_data = response_data["error"] 

143 error_message = error_data.get("message", "Unknown error") 

144 

145 choices = [] 

146 for choice_data in response_data.get("choices", []): 

147 message_data = choice_data.get("message", {}) 

148 

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 

179 

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) 

190 

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 ) 

200 

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 ) 

212 

213 def get_endpoint(self) -> str: 

214 return "/chat/completions" 

215 

216 

217class AnthropicAdapter(ProviderAdapter): 

218 """Adapter for Anthropic API.""" 

219 

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 

225 

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 }) 

235 

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] 

240 

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 } 

247 

248 # Add system message if present 

249 if system_message: 

250 anthropic_request["system"] = system_message 

251 

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 

267 

268 return anthropic_request 

269 

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 

279 

280 if "error" in response_data: 

281 success = False 

282 error_data = response_data["error"] 

283 error_message = error_data.get("message", "Unknown error") 

284 

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", "") 

293 

294 message = Message( 

295 role=MessageRole.ASSISTANT, 

296 content=content 

297 ) 

298 

299 choice = Choice( 

300 index=0, 

301 message=message, 

302 finish_reason=response_data.get("stop_reason") 

303 ) 

304 

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 ) 

314 

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 ) 

326 

327 def get_endpoint(self) -> str: 

328 return "/messages" 

329 

330 

331class OpenRouterAdapter(ProviderAdapter): 

332 """Adapter for OpenRouter API.""" 

333 

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 }) 

343 

344 model = request.model.strip() 

345 

346 # Build request 

347 openrouter_request: dict[str, Any] = { 

348 "model": model, 

349 "messages": messages, 

350 } 

351 

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 

371 

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} 

376 

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 } 

383 

384 return openrouter_request 

385 

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 

395 

396 if "error" in response_data: 

397 success = False 

398 error_data = response_data["error"] 

399 error_message = error_data.get("message", "Unknown error") 

400 

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) 

414 

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 ) 

424 

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 ) 

436 

437 def get_endpoint(self) -> str: 

438 return "/chat/completions" 

439 

440 

441# Provider adapter registry 

442PROVIDER_ADAPTERS = { 

443 Provider.OPENAI: OpenAIAdapter(), 

444 Provider.ANTHROPIC: AnthropicAdapter(), 

445 Provider.OPENROUTER: OpenRouterAdapter(), 

446} 

447 

448 

449def get_adapter(provider: Provider) -> ProviderAdapter: 

450 """Get the appropriate adapter for a provider.""" 

451 return PROVIDER_ADAPTERS[provider]