| """ |
| Base agent client — OpenAI-compatible wrapper around AMD's free Model API. |
| |
| Same architecture as v1 — the agent layer is intentionally model-agnostic |
| so it can be swapped to any OpenAI-compatible endpoint. SDXL-specific |
| prompt engineering lives in prompt_engineer.py. |
| """ |
| from __future__ import annotations |
|
|
| import logging |
| from dataclasses import dataclass |
| from typing import Any, Dict, List, Optional |
|
|
| from openai import OpenAI |
|
|
| from config.settings import settings |
|
|
| log = logging.getLogger(__name__) |
|
|
|
|
| @dataclass |
| class AgentResponse: |
| content: str |
| model: str |
| ok: bool |
| error: Optional[str] = None |
| raw: Optional[Dict[str, Any]] = None |
|
|
|
|
| class AgentClient: |
| """Thin wrapper around OpenAI SDK pointed at AMD's free API.""" |
|
|
| def __init__( |
| self, |
| model: Optional[str] = None, |
| fallback_model: Optional[str] = None, |
| api_key: Optional[str] = None, |
| base_url: Optional[str] = None, |
| temperature: float = 0.7, |
| max_tokens: int = 1024, |
| ) -> None: |
| self.model = model or settings.amd_agent_model |
| self.fallback_model = fallback_model or settings.amd_agent_fallback |
| self.temperature = temperature |
| self.max_tokens = max_tokens |
| self._api_key = api_key or settings.amd_api_key |
| self._base_url = base_url or settings.amd_base_url |
| self._client: Optional[OpenAI] = None |
|
|
| if not self._api_key: |
| log.warning( |
| "AMD_MODEL_API_KEY not set — agent layer disabled. " |
| "Core generation is unaffected." |
| ) |
|
|
| @property |
| def enabled(self) -> bool: |
| return bool(self._api_key) |
|
|
| def _get_client(self) -> OpenAI: |
| if self._client is None: |
| self._client = OpenAI(api_key=self._api_key, base_url=self._base_url) |
| return self._client |
|
|
| def chat( |
| self, |
| system_prompt: str, |
| user_prompt: str, |
| temperature: Optional[float] = None, |
| max_tokens: Optional[int] = None, |
| ) -> AgentResponse: |
| if not self.enabled: |
| return AgentResponse( |
| content="", model=self.model, ok=False, |
| error="Agent API key not configured", |
| ) |
|
|
| messages = [ |
| {"role": "system", "content": system_prompt}, |
| {"role": "user", "content": user_prompt}, |
| ] |
| return self._call_with_fallback(messages, temperature, max_tokens) |
|
|
| def _call_with_fallback( |
| self, |
| messages: List[Dict[str, str]], |
| temperature: Optional[float], |
| max_tokens: Optional[int], |
| ) -> AgentResponse: |
| models_to_try = [self.model] |
| if self.fallback_model and self.fallback_model != self.model: |
| models_to_try.append(self.fallback_model) |
|
|
| client = self._get_client() |
| last_error: Optional[str] = None |
|
|
| for model in models_to_try: |
| try: |
| resp = client.chat.completions.create( |
| model=model, |
| messages=messages, |
| temperature=temperature if temperature is not None else self.temperature, |
| max_tokens=max_tokens or self.max_tokens, |
| ) |
| content = resp.choices[0].message.content or "" |
| return AgentResponse( |
| content=content.strip(), |
| model=model, |
| ok=True, |
| raw={"usage": resp.usage.model_dump() if resp.usage else None}, |
| ) |
| except Exception as exc: |
| last_error = f"{type(exc).__name__}: {exc}" |
| log.warning("Agent call to %s failed: %s", model, last_error) |
| continue |
|
|
| return AgentResponse( |
| content="", model=self.model, ok=False, |
| error=last_error or "Unknown error", |
| ) |
|
|