Initial commit
This commit is contained in:
@@ -0,0 +1,108 @@
|
||||
"""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()
|
||||
}
|
||||
Reference in New Issue
Block a user