Dev2506's picture
Add files using upload-large-folder tool
15d68eb verified
Raw
History Blame Contribute Delete
3.88 kB
"""
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",
)