Files
cmhub/apps/ai/providers/base.py
T

151 lines
4.3 KiB
Python

from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any, Mapping, Protocol, Sequence
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]
@dataclass(frozen=True)
class MultimodalImage:
data: bytes
mime_type: str = "image/png"
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 analyze_images(
self,
prompt: str,
model: ResolvedModel,
*,
images: Sequence[MultimodalImage],
parameters: Mapping[str, Any] | None = None,
) -> TextGenerationResult:
...
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