Coverage for src / osiris_cli / client.py: 0%
602 statements
« prev ^ index » next coverage.py v7.13.0, created at 2025-12-31 05:01 +0200
« prev ^ index » next coverage.py v7.13.0, created at 2025-12-31 05:01 +0200
1import copy
2import httpx
3from openai import OpenAI, AsyncOpenAI, APIConnectionError, APITimeoutError
4from typing import Optional, Union, Any, List, Dict, AsyncGenerator, Iterable, Callable
5from .config import settings, KNOWN_PROVIDERS
6import sys
7import time
8from abc import ABC, abstractmethod
9from types import SimpleNamespace
11# Import our new systems
12from .logger import get_logger
13from .errors import (
14 APIError, APITimeoutError as OsirisAPITimeoutError,
15 APIRateLimitError, APIAuthenticationError,
16 APIInsufficientCreditsError, APIModelNotFoundError,
17 APIServerError, wrap_exception
18)
19from .resilience import (
20 resilient_api_call, get_circuit_breaker,
21 get_rate_limiter, fallback_manager
22)
23from .cost_tracker import cost_tracker
25logger = get_logger()
27# --- Core Types & Constants ---
29SAFE_FALLBACKS = {
30 "openrouter": "openai/gpt-oss-20b:free",
31 "openai": "gpt-3.5-turbo",
32 "groq": "llama3-8b-8192",
33 "perplexity": "sonar",
34 "deepseek": "deepseek-chat",
35 "gemini": "gemini-1.5-flash"
36}
39class AdapterContext:
40 def __init__(
41 self,
42 supports_tools_fn: Callable[[str, str], bool],
43 mark_tools_fn: Callable[[str, str, bool], None],
44 allows_tool_fn: Callable[[str], bool],
45 ):
46 self._supports_tools = supports_tools_fn
47 self._mark_tools = mark_tools_fn
48 self._allows_tool = allows_tool_fn
50 def supports_tools(self, provider: str, model: str) -> bool:
51 return self._supports_tools(provider, model)
53 def mark_tools(self, provider: str, model: str, supports: bool):
54 self._mark_tools(provider, model, supports)
56 def allows_tool(self, tool_name: str) -> bool:
57 return self._allows_tool(tool_name)
60class DefaultAdapterContext(AdapterContext):
61 def __init__(self):
62 super().__init__(lambda provider, model: True, lambda provider, model, supports: None, lambda name: True)
64def _is_connection_error(exc: Exception) -> bool:
65 if isinstance(exc, (APIConnectionError, APITimeoutError, httpx.ConnectError, httpx.ReadTimeout)):
66 return True
67 msg = str(exc).lower()
68 return "connection" in msg or "connect error" in msg or "network" in msg
70def _map_exception(provider: str, exc: Exception, model: Optional[str] = None) -> Exception:
71 error_str = str(exc).lower()
73 if "timeout" in error_str:
74 return OsirisAPITimeoutError(provider, settings.timeout)
75 if "rate limit" in error_str or "429" in error_str:
76 return APIRateLimitError(provider)
77 if "401" in error_str or "403" in error_str or "unauthorized" in error_str or "auth" in error_str:
78 return APIAuthenticationError(provider, 401)
79 if "402" in error_str or "payment required" in error_str or "insufficient" in error_str:
80 return APIInsufficientCreditsError(provider)
81 if "404" in error_str or "not found" in error_str:
82 return APIModelNotFoundError(provider, model or "unknown")
83 if any(code in error_str for code in ["500", "502", "503", "504"]):
84 return APIServerError(provider, 500)
86 return APIError(f"{provider} API error: {exc}")
89def _error_line(provider: str, exc: Exception, model: Optional[str] = None) -> str:
90 mapped = _map_exception(provider, exc, model)
92 if isinstance(mapped, APIModelNotFoundError):
93 return "Error: Model Not Found (404)"
94 if isinstance(mapped, APIAuthenticationError):
95 return "Error: Unauthorized (401)"
96 if isinstance(mapped, APIInsufficientCreditsError):
97 return "Error: Insufficient credits (402)"
98 if isinstance(mapped, APIRateLimitError):
99 return "Error: Rate limit exceeded (429)"
100 if isinstance(mapped, OsirisAPITimeoutError):
101 return f"Error: Request timed out after {settings.timeout}s"
102 if isinstance(mapped, APIServerError):
103 return "Error: Server error (5xx)"
105 return f"Error: {mapped}"
107# --- Adapters ---
109class BaseAdapter(ABC):
110 def __init__(
111 self,
112 provider: str,
113 client: Union[OpenAI, AsyncOpenAI],
114 adapter_context: Optional[AdapterContext] = None
115 ):
116 self.provider = provider
117 self.client = client
118 self.context = adapter_context or DefaultAdapterContext()
120 @abstractmethod
121 async def chat_stream(self, messages: List[Dict], model: str, temperature: float, tools: Optional[List] = None) -> AsyncGenerator[Any, None]:
122 pass
124 @abstractmethod
125 async def chat(self, messages: List[Dict], model: str, temperature: float, tools: Optional[List] = None) -> Any:
126 pass
128 def _get_headers(self) -> Dict:
129 if self.provider == "openrouter":
130 return {"HTTP-Referer": "https://theosirislabs.com", "X-Title": "Osiris CLI"}
131 return {}
133 def _format_tools(self, tools: Optional[List]) -> Optional[List]:
134 if not tools:
135 return None
137 formatted: List[Dict[str, Any]] = []
138 for tool in tools:
139 if not isinstance(tool, dict):
140 continue
141 function_data = None
142 if "function" in tool and isinstance(tool["function"], dict):
143 function_data = tool["function"]
144 name = (function_data or tool).get("name") or ""
145 if not name:
146 continue
147 description = (function_data or tool).get("description") or ""
148 parameters = (function_data or tool).get("parameters")
149 if not isinstance(parameters, dict):
150 parameters = {}
151 if not self.context.allows_tool(name):
152 continue
153 formatted.append({
154 "type": "function",
155 "function": {
156 "name": name,
157 "description": description,
158 "parameters": parameters,
159 }
160 })
162 return formatted or None
164class OpenAIChatAdapter(BaseAdapter):
165 """Standard OpenAI Chat Completions"""
166 async def chat_stream(self, messages: List[Dict], model: str, temperature: float, tools: Optional[List] = None) -> AsyncGenerator[Any, None]:
167 _temp = 1.0 if (model.startswith("o1-") or "reasoning" in model.lower()) else temperature
168 kwargs = self._build_chat_kwargs(model, messages, _temp, tools, stream=True)
170 try:
171 logger.debug(f"Starting streaming API call: {self.provider}/{model}")
172 start_time = time.time()
174 stream = await self._execute_with_openrouter_fallback(self.client.chat.completions.create, kwargs)
176 async for chunk in stream:
177 yield chunk
179 duration_ms = (time.time() - start_time) * 1000
180 logger.performance(f"stream_{self.provider}_{model}", duration_ms)
182 except Exception as e:
183 logger.error(f"Stream error: {e}", exc_info=True)
184 if _is_connection_error(e) and settings.offline_fallback:
185 raise
186 error_str = str(e).lower()
187 if tools and "v1/responses" in error_str:
188 fallback_model = "gpt-4o"
189 logger.warning(f"Model {model} requires Responses API; retrying with {fallback_model}")
190 kwargs["model"] = fallback_model
191 stream = await self.client.chat.completions.create(**kwargs)
192 async for chunk in stream:
193 yield chunk
194 return
195 if "tool calling is not supported" in str(e).lower() and tools:
196 logger.warning(f"Model {model} doesn't support tools, retrying without tools")
197 kwargs.pop("tools", None)
198 stream = await self.client.chat.completions.create(**kwargs)
199 async for chunk in stream:
200 yield chunk
201 else:
202 # Wrap error as a chunk
203 delta = SimpleNamespace(content=f"❌ {_error_line(self.provider, e, model)}", tool_calls=None)
204 yield SimpleNamespace(choices=[SimpleNamespace(delta=delta, index=0)])
206 async def chat(self, messages: List[Dict], model: str, temperature: float, tools: Optional[List] = None) -> Any:
207 _temp = 1.0 if (model.startswith("o1-") or "reasoning" in model.lower()) else temperature
208 kwargs = self._build_chat_kwargs(model, messages, _temp, tools, stream=False)
210 try:
211 logger.debug(f"Starting API call: {self.provider}/{model}")
212 start_time = time.time()
214 response = await self._execute_with_openrouter_fallback(self.client.chat.completions.create, kwargs)
216 duration_ms = (time.time() - start_time) * 1000
218 # Track cost if usage info available
219 if hasattr(response, 'usage') and response.usage:
220 usage = response.usage
221 input_tokens = getattr(usage, 'prompt_tokens', 0)
222 output_tokens = getattr(usage, 'completion_tokens', 0)
224 # Track the call
225 call = cost_tracker.track_call(
226 provider=self.provider,
227 model=model,
228 input_tokens=input_tokens,
229 output_tokens=output_tokens,
230 duration_ms=duration_ms,
231 session_id=settings.session_id
232 )
234 logger.info(
235 f"API call: {self.provider}/{model} - "
236 f"{call.total_tokens} tokens, ${call.total_cost:.4f}, {duration_ms:.0f}ms"
237 )
239 return response
241 except Exception as e:
242 logger.error(f"API call error: {e}", exc_info=True)
243 error_str = str(e).lower()
244 if tools and "v1/responses" in error_str:
245 fallback_model = "gpt-4o"
246 logger.warning(f"Model {model} requires Responses API; retrying with {fallback_model}")
247 kwargs["model"] = fallback_model
248 return await self.client.chat.completions.create(**kwargs)
249 if "tool calling is not supported" in str(e).lower() and tools:
250 logger.warning(f"Model {model} doesn't support tools, retrying without tools")
251 kwargs.pop("tools", None)
252 return await self.client.chat.completions.create(**kwargs)
253 raise self._convert_exception(e)
255 def _build_chat_kwargs(
256 self,
257 model: str,
258 messages: List[Dict],
259 temperature: float,
260 tools: Optional[List],
261 stream: bool
262 ) -> Dict[str, Any]:
263 kwargs = {
264 "model": model,
265 "messages": messages,
266 "temperature": temperature,
267 "extra_headers": self._get_headers(),
268 "timeout": settings.timeout,
269 }
270 if stream:
271 kwargs["stream"] = True
273 formatted_tools = self._format_tools(tools)
274 if formatted_tools and self.context.supports_tools(self.provider, model):
275 kwargs["tools"] = formatted_tools
276 kwargs["tool_choice"] = "auto"
278 return kwargs
280 async def _execute_with_openrouter_fallback(self, func: Callable[..., Any], kwargs: Dict[str, Any]) -> Any:
281 call_kwargs = dict(kwargs)
282 try:
283 result = await func(**call_kwargs)
284 except Exception as exc:
285 if self._should_retry_openrouter_tool_404(exc, call_kwargs):
286 model = call_kwargs.get("model")
287 self.context.mark_tools(self.provider, model, False)
288 clean_kwargs = self._strip_tool_fields(dict(call_kwargs))
289 logger.debug("OpenRouter fallback: retrying request without tools")
290 return await func(**clean_kwargs)
291 raise
292 response_model = getattr(result, "model", None)
293 self._cache_tool_support_models(call_kwargs.get("model"), response_model)
294 return result
296 def _should_retry_openrouter_tool_404(self, exc: Exception, kwargs: Dict[str, Any]) -> bool:
297 if self.provider != "openrouter":
298 return False
299 message = str(exc).lower()
300 if "no endpoints found that support tool use" not in message:
301 return False
302 return bool(kwargs.get("tools"))
304 def _strip_tool_fields(self, kwargs: Dict[str, Any]) -> Dict[str, Any]:
305 kwargs.pop("tools", None)
306 kwargs.pop("tool_choice", None)
307 return kwargs
309 def _cache_tool_support_models(self, requested_model: Optional[str], response_model: Optional[str]):
310 if not requested_model:
311 return
312 self.context.mark_tools(self.provider, requested_model, True)
313 if response_model and response_model != requested_model:
314 self.context.mark_tools(self.provider, response_model, True)
316 def _convert_exception(self, e: Exception) -> Exception:
317 """Convert OpenAI exceptions to Osiris exceptions"""
318 error_str = str(e).lower()
320 return _map_exception(self.provider, e)
322class OpenAIResponsesAdapter(BaseAdapter):
323 """OpenAI Beta Responses API"""
325 def _sanitize_responses_input(self, messages: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
326 sanitized: List[Dict[str, Any]] = []
328 for idx, message in enumerate(messages):
329 role = message.get("role", "unknown")
330 content = message.get("content")
331 if content is None:
332 is_placeholder = role == "tool" or bool(message.get("tool_calls"))
333 if is_placeholder:
334 logger.info(f"Sanitized Responses input: dropped placeholder message (role={role}, index={idx})")
335 continue
336 logger.info(f"Sanitized Responses input: replaced missing content (role={role}, index={idx})")
337 sanitized.append({**message, "content": ""})
338 continue
340 if isinstance(content, list):
341 valid_blocks = []
342 for block in content:
343 if isinstance(block, dict):
344 text = block.get("text")
345 if text:
346 valid_blocks.append({**block, "text": text})
347 elif isinstance(block, str) and block:
348 valid_blocks.append({"text": block})
349 if not valid_blocks:
350 logger.info(f"Sanitized Responses input: removed empty content blocks (role={role}, index={idx})")
351 continue
352 sanitized.append({**message, "content": valid_blocks})
353 continue
355 if isinstance(content, str):
356 sanitized.append(message)
357 continue
359 logger.info(f"Sanitized Responses input: coerced content to string (role={role}, index={idx})")
360 sanitized.append({**message, "content": str(content)})
362 return sanitized
364 def _build_tool_call_delta(self, tool_call: Any) -> Optional[SimpleNamespace]:
365 if tool_call is None:
366 return None
368 if isinstance(tool_call, dict):
369 name = tool_call.get("name") or tool_call.get("function", {}).get("name")
370 arguments = tool_call.get("arguments") or tool_call.get("function", {}).get("arguments")
371 call_id = tool_call.get("id") or tool_call.get("call_id", "")
372 index = tool_call.get("index", 0)
373 else:
374 name = getattr(tool_call, "name", None) or getattr(getattr(tool_call, "function", None), "name", None)
375 arguments = getattr(tool_call, "arguments", None) or getattr(getattr(tool_call, "function", None), "arguments", None)
376 call_id = getattr(tool_call, "id", None) or getattr(tool_call, "call_id", "")
377 index = getattr(tool_call, "index", 0)
379 if not name:
380 return None
382 fn = SimpleNamespace(name=name, arguments=arguments or "")
383 return SimpleNamespace(index=index or 0, id=call_id or "", function=fn)
385 def _extract_tool_call(self, chunk: Any) -> List[SimpleNamespace]:
386 tool_calls: List[SimpleNamespace] = []
388 chunk_type = getattr(chunk, "type", "")
389 if chunk_type in {"response.output_item.added", "response.output_item.done"}:
390 item = getattr(chunk, "item", None)
391 if item is None and hasattr(chunk, "data"):
392 data = chunk.data
393 if isinstance(data, dict):
394 item = data.get("item")
395 elif hasattr(data, "item"):
396 item = data.item
398 if item is not None:
399 if isinstance(item, dict) and item.get("type") in {"tool_call", "function_call"}:
400 delta = self._build_tool_call_delta(item)
401 if delta:
402 tool_calls.append(delta)
403 return tool_calls
404 if hasattr(item, "type") and item.type in {"tool_call", "function_call"}:
405 delta = self._build_tool_call_delta(item)
406 if delta:
407 tool_calls.append(delta)
408 return tool_calls
410 if hasattr(chunk, "tool_calls") and chunk.tool_calls:
411 for tc in chunk.tool_calls:
412 delta = self._build_tool_call_delta(tc)
413 if delta:
414 tool_calls.append(delta)
415 return tool_calls
417 if hasattr(chunk, "tool_call"):
418 delta = self._build_tool_call_delta(chunk.tool_call)
419 if delta:
420 tool_calls.append(delta)
421 return tool_calls
423 if hasattr(chunk, "data"):
424 data = chunk.data
425 if isinstance(data, dict) and "tool_call" in data:
426 delta = self._build_tool_call_delta(data.get("tool_call"))
427 if delta:
428 tool_calls.append(delta)
429 elif hasattr(data, "tool_call"):
430 delta = self._build_tool_call_delta(data.tool_call)
431 if delta:
432 tool_calls.append(delta)
433 if tool_calls:
434 return tool_calls
436 if hasattr(chunk, "item"):
437 item = chunk.item
438 if isinstance(item, dict) and "tool_call" in item:
439 delta = self._build_tool_call_delta(item.get("tool_call"))
440 if delta:
441 tool_calls.append(delta)
442 elif hasattr(item, "tool_call"):
443 delta = self._build_tool_call_delta(item.tool_call)
444 if delta:
445 tool_calls.append(delta)
447 return tool_calls
449 async def chat_stream(self, messages: List[Dict], model: str, temperature: float, tools: Optional[List] = None) -> AsyncGenerator[Any, None]:
450 instructions = next((m["content"] for m in messages if m["role"] == "system"), "You are a helpful assistant.")
451 user_input = [m for m in messages if m["role"] != "system"]
452 sanitized_input = self._sanitize_responses_input(user_input)
453 kwargs = {
454 "model": model, "instructions": instructions, "input": sanitized_input, "stream": True,
455 "extra_headers": self._get_headers(), "timeout": settings.timeout
456 }
457 formatted_tools = self._format_tools(tools)
458 if formatted_tools:
459 kwargs["tools"] = formatted_tools
460 kwargs["tool_choice"] = "auto"
462 last_emitted: Optional[str] = None
463 accumulated: str = ""
465 try:
466 logger.debug(f"Starting Responses API stream: {self.provider}/{model}")
467 stream = await self.client.responses.create(**kwargs)
468 async for chunk in stream:
469 if hasattr(chunk, "type") and "tool_call" in str(chunk.type):
470 tool_calls = self._extract_tool_call(chunk)
471 if tool_calls:
472 delta = SimpleNamespace(content=None, tool_calls=tool_calls)
473 yield SimpleNamespace(choices=[SimpleNamespace(delta=delta, index=0)])
474 continue
476 tool_calls = self._extract_tool_call(chunk)
477 if tool_calls:
478 delta = SimpleNamespace(content=None, tool_calls=tool_calls)
479 yield SimpleNamespace(choices=[SimpleNamespace(delta=delta, index=0)])
480 continue
482 # Map chunk to standard delta.content if possible
483 if hasattr(chunk, 'choices') and chunk.choices:
484 yield chunk
485 elif hasattr(chunk, 'type') and chunk.type == "content_part" and hasattr(chunk, 'part'):
486 content = getattr(chunk.part, 'text', '')
487 if content:
488 if content.startswith(accumulated):
489 delta_text = content[len(accumulated):]
490 else:
491 delta_text = content
492 if delta_text and delta_text != last_emitted:
493 last_emitted = delta_text
494 accumulated += delta_text
495 delta = SimpleNamespace(content=delta_text, tool_calls=None)
496 yield SimpleNamespace(choices=[SimpleNamespace(delta=delta, index=0)])
497 elif hasattr(chunk, 'type'):
498 content = ""
499 if hasattr(chunk, 'delta'):
500 content = chunk.delta or ""
501 elif hasattr(chunk, 'text'):
502 content = chunk.text or ""
503 elif hasattr(chunk, 'content'):
504 content = chunk.content or ""
505 if content:
506 if content.startswith(accumulated):
507 delta_text = content[len(accumulated):]
508 else:
509 delta_text = content
510 if delta_text and delta_text != last_emitted:
511 last_emitted = delta_text
512 accumulated += delta_text
513 delta = SimpleNamespace(content=delta_text, tool_calls=None)
514 yield SimpleNamespace(choices=[SimpleNamespace(delta=delta, index=0)])
515 elif hasattr(chunk, 'status') and chunk.status == "failed":
516 delta = SimpleNamespace(content=f"❌ Model status failed.", tool_calls=None)
517 yield SimpleNamespace(choices=[SimpleNamespace(delta=delta, index=0)])
518 except Exception as e:
519 logger.error(f"Responses API stream error: {e}", exc_info=True)
520 if _is_connection_error(e) and settings.offline_fallback:
521 raise
522 delta = SimpleNamespace(content=f"❌ Responses API Error: {e}", tool_calls=None)
523 yield SimpleNamespace(choices=[SimpleNamespace(delta=delta, index=0)])
525 async def chat(self, messages: List[Dict], model: str, temperature: float, tools: Optional[List] = None) -> Any:
526 instructions = next((m["content"] for m in messages if m["role"] == "system"), "You are a helpful assistant.")
527 user_input = [m for m in messages if m["role"] != "system"]
528 sanitized_input = self._sanitize_responses_input(user_input)
529 kwargs = {
530 "model": model, "instructions": instructions, "input": sanitized_input,
531 "extra_headers": self._get_headers(), "timeout": settings.timeout
532 }
533 formatted_tools = self._format_tools(tools)
534 if formatted_tools and self.context.supports_tools(self.provider, model):
535 kwargs["tools"] = formatted_tools
536 kwargs["tool_choice"] = "auto"
538 logger.debug(f"Starting Responses API call: {self.provider}/{model}")
539 start_time = time.time()
540 try:
541 res = await self.client.responses.create(**kwargs)
542 duration_ms = (time.time() - start_time) * 1000
543 logger.performance(f"responses_{self.provider}_{model}", duration_ms)
544 except Exception as e:
545 logger.error(f"Responses API call error: {e}", exc_info=True)
546 if _is_connection_error(e) and settings.offline_fallback:
547 raise
548 raise _map_exception(self.provider, e, model)
550 if hasattr(res, 'choices'): return res
552 # Shim for Responses object to look like Completion object
553 content = ""
554 if hasattr(res, 'output') and res.output:
555 for part in res.output:
556 if hasattr(part, 'text'): content += part.text
558 msg = SimpleNamespace(content=content, role="assistant", tool_calls=getattr(res, 'tool_calls', None))
559 return SimpleNamespace(choices=[SimpleNamespace(message=msg, finish_reason=getattr(res, 'status', 'stop'))])
561# --- Factory ---
563def get_adapter(
564 provider: str,
565 model: str,
566 client: Union[OpenAI, AsyncOpenAI],
567 adapter_context: Optional[AdapterContext] = None
568) -> BaseAdapter:
569 context = adapter_context or DefaultAdapterContext()
570 if provider == "openai":
571 m = model.lower()
572 if any(x in m for x in ["gpt-5", "o3", "codex-mini"]):
573 return OpenAIResponsesAdapter(provider, client, context)
574 return OpenAIChatAdapter(provider, client, context)
576# --- Client Classes ---
578class AsyncLLMClient:
579 def __init__(self):
580 self.provider = settings.provider
581 self.client = self._setup_client()
582 logger.info(f"Initialized AsyncLLMClient with provider: {self.provider}")
584 def _adapter_context(self) -> AdapterContext:
585 from .session import session
586 return AdapterContext(
587 lambda provider, model: session.provider_supports_tools(provider, model),
588 lambda provider, model, supports: session.mark_provider_tool_support(provider, model, supports),
589 lambda tool_name: session.is_tool_allowed(tool_name),
590 )
592 def _setup_offline_client(self) -> Optional[AsyncOpenAI]:
593 if not settings.offline_fallback:
594 return None
595 provider = settings.offline_provider
596 provider_data = KNOWN_PROVIDERS.get(provider)
597 if not provider_data:
598 return None
599 api_key = "ollama"
600 if provider_data.get("api_key_name"):
601 val = getattr(settings, provider_data["api_key_name"], None)
602 if val:
603 api_key = val.get_secret_value() if hasattr(val, "get_secret_value") else val
604 base_url = provider_data["base_url"]
605 if provider == "ollama":
606 base_url = settings.ollama_base_url
607 return AsyncOpenAI(base_url=base_url, api_key=api_key)
609 def _setup_client(self) -> AsyncOpenAI:
610 provider_data = KNOWN_PROVIDERS.get(self.provider)
611 if not provider_data:
612 logger.error(f"Unknown provider: {self.provider}")
613 return AsyncOpenAI(api_key="missing")
615 api_key = os.getenv("OLLAMA_API_KEY", "ollama")
616 if provider_data["api_key_name"]:
617 val = getattr(settings, provider_data["api_key_name"], None)
618 if val:
619 api_key = val.get_secret_value() if hasattr(val, "get_secret_value") else val
620 else:
621 logger.warning(f"API key not configured for {self.provider}")
622 api_key = "needs_configuration"
624 base_url = settings.provider_base_url or provider_data["base_url"]
625 logger.debug(f"Client setup: {self.provider} at {base_url}")
627 return AsyncOpenAI(base_url=base_url, api_key=api_key)
629 def reinitialize(self):
630 logger.info(f"Reinitializing client for provider: {settings.provider}")
631 self.provider = settings.provider
632 self.client = self._setup_client()
634 @resilient_api_call(max_attempts=3, min_wait=1.0, max_wait=10.0)
635 async def chat(self, messages: List[Dict], model: Optional[str] = None, temperature: Optional[float] = None, tools: Optional[List] = None) -> Any:
636 _model = model or settings.default_model
637 _temp = temperature if temperature is not None else settings.temperature
639 rate_limiter = get_rate_limiter(self.provider)
640 await rate_limiter.acquire()
642 circuit_breaker = get_circuit_breaker(self.provider)
644 adapter_context = self._adapter_context()
645 try:
646 adapter = get_adapter(self.provider, _model, self.client, adapter_context)
647 response = await adapter.chat(messages, _model, _temp, tools)
648 circuit_breaker._on_success()
649 return response
650 except Exception as e:
651 logger.error(f"Async API call failed: {e}", exc_info=True)
652 circuit_breaker._on_failure()
653 raise _map_exception(self.provider, e, _model)
655 async def chat_stream(self, messages: List[Dict], model: Optional[str] = None, temperature: Optional[float] = None) -> AsyncGenerator[str, None]:
656 _model = model or settings.default_model
657 _temp = temperature if temperature is not None else settings.temperature
659 # Check rate limit
660 rate_limiter = get_rate_limiter(self.provider)
661 await rate_limiter.acquire()
663 # Use circuit breaker
664 circuit_breaker = get_circuit_breaker(self.provider)
666 adapter_context = self._adapter_context()
667 try:
668 adapter = get_adapter(self.provider, _model, self.client, adapter_context)
669 async for chunk in adapter.chat_stream(messages, _model, _temp):
670 if chunk.choices and chunk.choices[0].delta.content:
671 yield chunk.choices[0].delta.content
673 circuit_breaker._on_success()
675 except Exception as e:
676 logger.error(f"Stream failed: {e}", exc_info=True)
677 circuit_breaker._on_failure()
679 # Try fallback
680 fallback_provider = fallback_manager.get_fallback_provider(self.provider)
681 if fallback_provider:
682 logger.warning(f"Attempting fallback to {fallback_provider}")
683 # Would need to switch provider and retry here
685 raise _map_exception(self.provider, e, _model)
687 @resilient_api_call(max_attempts=3, min_wait=1.0, max_wait=10.0)
688 async def create_completion(self, messages: List[Dict], model: Optional[str] = None, temperature: Optional[float] = None, tools: Optional[List] = None, stream: bool = False) -> Any:
689 _model = model or settings.default_model
690 _temp = temperature if temperature is not None else settings.temperature
692 # Check rate limit
693 rate_limiter = get_rate_limiter(self.provider)
694 await rate_limiter.acquire()
696 # Force Chat Completions for tool calls to avoid Responses input incompatibilities.
697 adapter_context = self._adapter_context()
698 if tools and self.provider == "openai":
699 adapter = OpenAIChatAdapter(self.provider, self.client, adapter_context)
700 else:
701 adapter = get_adapter(self.provider, _model, self.client, adapter_context)
703 if stream:
704 async def stream_wrapper():
705 try:
706 async for chunk in adapter.chat_stream(messages, _model, _temp, tools):
707 yield chunk
708 except Exception as e:
709 if _is_connection_error(e) and settings.offline_fallback and self.provider != settings.offline_provider:
710 offline_client = self._setup_offline_client()
711 if offline_client:
712 offline_model = settings.offline_model or _model
713 logger.warning(f"Falling back to offline provider {settings.offline_provider}/{offline_model}")
714 offline_context = self._adapter_context()
715 offline_adapter = get_adapter(settings.offline_provider, offline_model, offline_client, offline_context)
716 async for chunk in offline_adapter.chat_stream(messages, offline_model, _temp, tools):
717 yield chunk
718 return
719 raise
720 return stream_wrapper()
721 try:
722 return await adapter.chat(messages, _model, _temp, tools)
723 except Exception as e:
724 if _is_connection_error(e) and settings.offline_fallback and self.provider != settings.offline_provider:
725 offline_client = self._setup_offline_client()
726 if offline_client:
727 offline_model = settings.offline_model or _model
728 logger.warning(f"Falling back to offline provider {settings.offline_provider}/{offline_model}")
729 offline_context = self._adapter_context()
730 offline_adapter = get_adapter(settings.offline_provider, offline_model, offline_client, offline_context)
731 return await offline_adapter.chat(messages, offline_model, _temp, tools)
732 raise
734class LLMClient:
735 def __init__(self):
736 self.provider = settings.provider
737 self.client = self._setup_client()
738 logger.info(f"Initialized LLMClient with provider: {self.provider}")
740 def _setup_client(self) -> OpenAI:
741 provider_data = KNOWN_PROVIDERS.get(self.provider)
742 if not provider_data: return OpenAI(api_key="missing")
743 api_key = "ollama"
744 if provider_data["api_key_name"]:
745 val = getattr(settings, provider_data["api_key_name"], None)
746 if val:
747 api_key = val.get_secret_value() if hasattr(val, "get_secret_value") else val
748 base_url = settings.provider_base_url or provider_data["base_url"]
749 return OpenAI(base_url=base_url, api_key=api_key)
751 def chat(self, messages: List[Dict], model: Optional[str] = None, temperature: Optional[float] = None, tools: Optional[List] = None) -> Any:
752 _model = model or settings.default_model
753 use_responses = (_model.lower() and any(x in _model.lower() for x in ["gpt-5", "o3", "codex-mini"])) and self.provider == "openai"
755 logger.debug(f"Sync API call: {self.provider}/{_model}")
756 start_time = time.time()
758 try:
759 if use_responses:
760 instructions = next((m["content"] for m in messages if m["role"] == "system"), None)
761 user_input = [m for m in messages if m["role"] != "system"]
762 response = self.client.responses.create(model=_model, instructions=instructions, input=user_input, tools=tools)
763 else:
764 response = self.client.chat.completions.create(model=_model, messages=messages, temperature=temperature or 0.4, tools=tools)
766 duration_ms = (time.time() - start_time) * 1000
767 logger.performance(f"sync_{self.provider}_{_model}", duration_ms)
769 return response
771 except Exception as e:
772 logger.error(f"Sync API call failed: {e}", exc_info=True)
773 raise _map_exception(self.provider, e, _model)
775# --- Global Helpers ---
777import os
778# Note: clients initialized on first use to avoid import-time errors
779async_client = None
780client = None
782def get_async_client():
783 global async_client
784 if async_client is None:
785 async_client = AsyncLLMClient()
786 return async_client
788def get_client():
789 global client
790 if client is None:
791 client = LLMClient()
792 return client
794_MODEL_CACHE: Dict[str, tuple] = {}
795_CACHE_TTL = 300.0
797async def fetch_models_dynamic(provider: str, api_key: Optional[str] = None) -> List[str]:
798 cache_key = f"{provider}:{api_key}"
799 now = time.monotonic()
800 if cache_key in _MODEL_CACHE:
801 ms, ts = _MODEL_CACHE[cache_key]
802 if now - ts < _CACHE_TTL: return ms
804 provider_data = KNOWN_PROVIDERS.get(provider)
805 if not provider_data: return []
806 base_url = provider_data["base_url"]
807 real_key = api_key or "missing"
808 if provider == "ollama": real_key = "ollama"
810 models = []
811 try:
812 logger.debug(f"Fetching models for {provider}")
813 if provider == "openrouter":
814 async with httpx.AsyncClient() as h:
815 r = await h.get("https://openrouter.ai/api/v1/models")
816 models = sorted([m["id"] for m in r.json()["data"]])
817 else:
818 temp = AsyncOpenAI(base_url=base_url, api_key=real_key)
819 resp = await temp.models.list()
820 models = sorted([m.id for m in resp.data])
822 logger.info(f"Fetched {len(models)} models for {provider}")
823 except Exception as e:
824 logger.warning(f"Failed to fetch models for {provider}: {e}")
825 models = provider_data.get("models", [])
827 if models: _MODEL_CACHE[cache_key] = (models, now)
828 return models