109 lines
3.4 KiB
Python
109 lines
3.4 KiB
Python
"""Project-level routing for provider-qualified OpenAI-compatible models."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from dataclasses import dataclass
|
|
import json
|
|
import os
|
|
from pathlib import Path
|
|
|
|
from dotenv import load_dotenv
|
|
|
|
|
|
PROJECT_ROOT = Path(__file__).resolve().parent.parent
|
|
ENV_FILE = PROJECT_ROOT / ".env"
|
|
ROUTES_FILE = PROJECT_ROOT / "provider_routes.json"
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class ModelReference:
|
|
provider: str
|
|
model_id: str
|
|
|
|
@property
|
|
def value(self) -> str:
|
|
return f"{self.provider}/{self.model_id}"
|
|
|
|
@property
|
|
def slug(self) -> str:
|
|
return self.value.replace("/", "_")
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class ProviderConfig:
|
|
name: str
|
|
label: str
|
|
url_env: str
|
|
key_env: str
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class ModelRoute:
|
|
reference: ModelReference
|
|
config: ProviderConfig
|
|
url: str
|
|
api_key: str
|
|
|
|
|
|
def provider_configs() -> dict[str, ProviderConfig]:
|
|
try:
|
|
raw = json.loads(ROUTES_FILE.read_text(encoding="utf-8"))
|
|
except (OSError, json.JSONDecodeError) as exc:
|
|
raise RuntimeError(f"cannot load provider routes from {ROUTES_FILE}: {exc}") from exc
|
|
if not isinstance(raw, dict) or not raw:
|
|
raise RuntimeError(f"provider routes must be a non-empty JSON object: {ROUTES_FILE}")
|
|
routes: dict[str, ProviderConfig] = {}
|
|
for name, value in raw.items():
|
|
if not isinstance(name, str) or not name or not isinstance(value, dict):
|
|
raise RuntimeError(f"invalid provider route in {ROUTES_FILE}")
|
|
try:
|
|
routes[name] = ProviderConfig(
|
|
name=name,
|
|
label=str(value.get("label") or name),
|
|
url_env=str(value["url_env"]),
|
|
key_env=str(value["key_env"]),
|
|
)
|
|
except KeyError as exc:
|
|
raise RuntimeError(f"provider '{name}' is missing {exc.args[0]} in {ROUTES_FILE}") from exc
|
|
return routes
|
|
|
|
|
|
def parse_model_reference(value: str) -> ModelReference:
|
|
raw = value.strip().strip("/")
|
|
provider, separator, model_id = raw.partition("/")
|
|
if not separator or not provider or not model_id:
|
|
raise ValueError("model must use provider/model-id format, for example opencode/deepseek-v4-pro")
|
|
if any(part in {"", ".", ".."} for part in raw.split("/")):
|
|
raise ValueError("model must not contain empty or relative path segments")
|
|
routes = provider_configs()
|
|
if provider not in routes:
|
|
allowed = ", ".join(sorted(routes))
|
|
raise ValueError(f"unsupported provider '{provider}'; configured providers: {allowed}")
|
|
return ModelReference(provider, model_id)
|
|
|
|
|
|
def resolve_model_route(
|
|
model: str | ModelReference,
|
|
*,
|
|
require_credentials: bool = True,
|
|
) -> ModelRoute | None:
|
|
reference = parse_model_reference(model) if isinstance(model, str) else model
|
|
config = provider_configs()[reference.provider]
|
|
load_dotenv(ENV_FILE, override=False)
|
|
url = os.getenv(config.url_env, "").strip()
|
|
api_key = os.getenv(config.key_env, "").strip()
|
|
if not url or not api_key:
|
|
if not require_credentials:
|
|
return None
|
|
raise RuntimeError(
|
|
f"provider '{reference.provider}' requires {config.url_env} and {config.key_env} in {ENV_FILE}"
|
|
)
|
|
return ModelRoute(reference, config, url, api_key)
|
|
|
|
|
|
def provider_environment() -> dict[str, tuple[str, str]]:
|
|
return {
|
|
name: (config.url_env, config.key_env)
|
|
for name, config in provider_configs().items()
|
|
}
|