135 lines
3.9 KiB
Python
135 lines
3.9 KiB
Python
from __future__ import annotations
|
|
|
|
from dataclasses import dataclass, field
|
|
from typing import Any, Mapping, Protocol
|
|
|
|
|
|
class AiProviderError(RuntimeError):
|
|
"""Base error for AI provider failures."""
|
|
|
|
|
|
class AiProviderConfigError(AiProviderError, ValueError):
|
|
"""Raised when model/provider configuration is invalid."""
|
|
|
|
|
|
class AiCapabilityError(AiProviderError):
|
|
"""Raised when a provider cannot perform the requested capability."""
|
|
|
|
|
|
class AiResponseParseError(AiProviderError):
|
|
"""Raised when a provider response cannot be parsed into the expected result."""
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class ResolvedModel:
|
|
"""Runtime model config resolved from future AiModel/ModelAlias records."""
|
|
|
|
name: str
|
|
url: str
|
|
model: str
|
|
api_key: str
|
|
api_type: str = "auto"
|
|
timeout_seconds: int = 0
|
|
connect_timeout_seconds: int = 30
|
|
extra_body: dict[str, Any] = field(default_factory=dict)
|
|
capabilities: frozenset[str] = field(default_factory=frozenset)
|
|
|
|
@classmethod
|
|
def from_mapping(cls, data: Mapping[str, Any]) -> "ResolvedModel":
|
|
raw_capabilities = data.get("capabilities") or ()
|
|
if isinstance(raw_capabilities, str):
|
|
capabilities = frozenset(
|
|
item.strip() for item in raw_capabilities.split(",") if item.strip()
|
|
)
|
|
elif isinstance(raw_capabilities, Mapping):
|
|
capabilities = frozenset(
|
|
str(key)
|
|
for key, value in raw_capabilities.items()
|
|
if key in {"text", "image", "vision"} and bool(value)
|
|
)
|
|
else:
|
|
capabilities = frozenset(str(item) for item in raw_capabilities)
|
|
|
|
return cls(
|
|
name=str(data.get("name") or data.get("model") or ""),
|
|
url=str(data.get("url") or ""),
|
|
model=str(data.get("model") or ""),
|
|
api_key=str(data.get("api_key") or ""),
|
|
api_type=str(data.get("api_type") or "auto"),
|
|
timeout_seconds=_to_int(data.get("timeout_seconds"), 0),
|
|
connect_timeout_seconds=_to_int(data.get("connect_timeout_seconds"), 30),
|
|
extra_body=dict(data.get("extra_body") or {}),
|
|
capabilities=capabilities,
|
|
)
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class TextGenerationResult:
|
|
text: str
|
|
titles: tuple[str, ...]
|
|
model_used: str
|
|
raw: Mapping[str, Any]
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class ImageGenerationResult:
|
|
image: bytes
|
|
model_used: str
|
|
raw: Mapping[str, Any]
|
|
|
|
|
|
class Provider(Protocol):
|
|
def capabilities(self) -> set[str]:
|
|
...
|
|
|
|
def generate_text(
|
|
self,
|
|
prompt: str,
|
|
model: ResolvedModel,
|
|
*,
|
|
image: bytes | None = None,
|
|
image_mime_type: str = "image/png",
|
|
resolution: str = "1K",
|
|
parameters: Mapping[str, Any] | None = None,
|
|
) -> TextGenerationResult:
|
|
...
|
|
|
|
def generate_image(
|
|
self,
|
|
prompt: str,
|
|
model: ResolvedModel,
|
|
*,
|
|
image: bytes | None = None,
|
|
image_mime_type: str = "image/png",
|
|
image_filename: str = "image.png",
|
|
resolution: str = "1K",
|
|
aspect_ratio: str = "1:1",
|
|
parameters: Mapping[str, Any] | None = None,
|
|
) -> ImageGenerationResult:
|
|
...
|
|
|
|
|
|
def validate_model_config(model: ResolvedModel) -> None:
|
|
errors = []
|
|
if not model.url.strip():
|
|
errors.append("missing url")
|
|
if not model.model.strip():
|
|
errors.append("missing model")
|
|
if not model.api_key.strip():
|
|
errors.append("missing api_key")
|
|
if model.timeout_seconds < 0:
|
|
errors.append("timeout_seconds must be >= 0")
|
|
if model.connect_timeout_seconds <= 0:
|
|
errors.append("connect_timeout_seconds must be > 0")
|
|
if errors:
|
|
raise AiProviderConfigError("; ".join(errors))
|
|
|
|
|
|
def _to_int(value: Any, default: int) -> int:
|
|
if value is None or value == "":
|
|
return default
|
|
try:
|
|
return int(value)
|
|
except (TypeError, ValueError):
|
|
return default
|