# retoor <retoor@molodetz.nl>
from __future__ import annotations
import logging
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Optional
from pydantic import BaseModel, Field, field_validator
from devplacepy.database import (
bump_cache_version,
db,
get_table,
sync_local_cache,
)
logger = logging.getLogger(__name__)
PROVIDERS_TABLE = "gateway_providers"
MODELS_TABLE = "gateway_models"
CACHE_NAME = "gateway_routing"
KINDS = ("chat", "embed")
_ROUTING_CACHE: dict = {}
def _now() -> str:
return datetime.now(timezone.utc).isoformat()
def _as_bool(value) -> bool:
if isinstance(value, bool):
return value
if isinstance(value, (int, float)):
return value != 0
return str(value).strip() in ("1", "true", "True", "yes", "on")
def _embed_url_from_base(base_url: str) -> str:
base_url = (base_url or "").strip()
if not base_url:
return ""
if base_url.endswith("/chat/completions"):
return base_url[: -len("/chat/completions")] + "/embeddings"
return base_url
def ensure_tables() -> None:
db.query(
"CREATE TABLE IF NOT EXISTS "
+ PROVIDERS_TABLE
+ " (id INTEGER PRIMARY KEY, name TEXT, base_url TEXT, api_key TEXT, "
"is_active INTEGER DEFAULT 1, created_at TEXT, updated_at TEXT)"
)
db.query(
"CREATE TABLE IF NOT EXISTS "
+ MODELS_TABLE
+ " (id INTEGER PRIMARY KEY, source_model TEXT, provider TEXT, "
"target_model TEXT, kind TEXT DEFAULT 'chat', vision_provider TEXT, "
"vision_model TEXT, context_window INTEGER DEFAULT 0, "
"price_cache_hit_per_m REAL DEFAULT 0, price_cache_miss_per_m REAL DEFAULT 0, "
"price_output_per_m REAL DEFAULT 0, price_input_per_m REAL DEFAULT 0, "
"is_active INTEGER DEFAULT 1, created_at TEXT, updated_at TEXT)"
)
try:
db.query(
"CREATE UNIQUE INDEX IF NOT EXISTS idx_gateway_providers_name ON "
+ PROVIDERS_TABLE
+ " (name)"
)
db.query(
"CREATE UNIQUE INDEX IF NOT EXISTS idx_gateway_models_source ON "
+ MODELS_TABLE
+ " (source_model)"
)
db.query(
"CREATE INDEX IF NOT EXISTS idx_gateway_models_kind ON "
+ MODELS_TABLE
+ " (kind)"
)
except Exception as exc:
logger.warning("gateway routing index creation failed: %s", exc)
class ProviderIn(BaseModel):
name: str = Field(min_length=1, max_length=64)
base_url: str = Field(default="", max_length=500)
api_key: str = Field(default="", max_length=400)
is_active: bool = True
@field_validator("name")
@classmethod
def _clean_name(cls, value: str) -> str:
value = value.strip().lower()
if not value:
raise ValueError("Provider name is required")
if not all(c.isalnum() or c in "-_" for c in value):
raise ValueError("Provider name allows letters, numbers, hyphen, underscore")
return value
@field_validator("base_url")
@classmethod
def _clean_url(cls, value: str) -> str:
value = (value or "").strip()
if value and not (value.startswith("http://") or value.startswith("https://")):
raise ValueError("Base URL must be a http(s) URL")
return value
class ModelRouteIn(BaseModel):
source_model: str = Field(min_length=1, max_length=128)
provider: str = Field(default="", max_length=64)
target_model: str = Field(min_length=1, max_length=128)
kind: str = "chat"
vision_provider: str = Field(default="", max_length=64)
vision_model: str = Field(default="", max_length=128)
context_window: int = Field(default=0, ge=0, le=100_000_000)
price_cache_hit_per_m: float = Field(default=0.0, ge=0)
price_cache_miss_per_m: float = Field(default=0.0, ge=0)
price_output_per_m: float = Field(default=0.0, ge=0)
price_input_per_m: float = Field(default=0.0, ge=0)
is_active: bool = True
@field_validator("source_model", "target_model")
@classmethod
def _clean_model(cls, value: str) -> str:
value = (value or "").strip()
if not value:
raise ValueError("Model name is required")
return value
@field_validator("provider", "vision_provider", "vision_model")
@classmethod
def _strip(cls, value: str) -> str:
return (value or "").strip()
@field_validator("kind")
@classmethod
def _clean_kind(cls, value: str) -> str:
value = (value or "chat").strip().lower()
if value not in KINDS:
raise ValueError("Kind must be 'chat' or 'embed'")
return value
@dataclass(frozen=True)
class ModelRoute:
source_model: str
provider: str
target_model: str
kind: str
vision_provider: str
vision_model: str
context_window: int
price_cache_hit_per_m: float
price_cache_miss_per_m: float
price_output_per_m: float
price_input_per_m: float
is_active: bool
def _route_from_row(row: dict) -> ModelRoute:
return ModelRoute(
source_model=str(row.get("source_model") or ""),
provider=str(row.get("provider") or ""),
target_model=str(row.get("target_model") or ""),
kind=str(row.get("kind") or "chat"),
vision_provider=str(row.get("vision_provider") or ""),
vision_model=str(row.get("vision_model") or ""),
context_window=int(row.get("context_window") or 0),
price_cache_hit_per_m=float(row.get("price_cache_hit_per_m") or 0.0),
price_cache_miss_per_m=float(row.get("price_cache_miss_per_m") or 0.0),
price_output_per_m=float(row.get("price_output_per_m") or 0.0),
price_input_per_m=float(row.get("price_input_per_m") or 0.0),
is_active=_as_bool(row.get("is_active", 1)),
)
def _load() -> dict:
sync_local_cache(CACHE_NAME, _ROUTING_CACHE)
if "providers" not in _ROUTING_CACHE:
providers: dict = {}
models: dict = {}
try:
if PROVIDERS_TABLE in db.tables:
for row in get_table(PROVIDERS_TABLE).all():
name = str(row.get("name") or "").strip().lower()
if name:
providers[name] = {
"name": name,
"base_url": str(row.get("base_url") or ""),
"api_key": str(row.get("api_key") or ""),
"is_active": _as_bool(row.get("is_active", 1)),
}
if MODELS_TABLE in db.tables:
for row in get_table(MODELS_TABLE).all():
source = str(row.get("source_model") or "").strip()
if source:
models[source] = _route_from_row(row)
except Exception as exc:
logger.warning("gateway routing load failed: %s", exc)
_ROUTING_CACHE["providers"] = providers
_ROUTING_CACHE["models"] = models
return _ROUTING_CACHE
class ProviderStore:
def list(self) -> list[dict]:
return sorted(_load()["providers"].values(), key=lambda p: p["name"])
def get(self, name: str) -> Optional[dict]:
if not name:
return None
return _load()["providers"].get(name.strip().lower())
def names(self) -> list[str]:
return sorted(_load()["providers"].keys())
def count(self) -> int:
return len(_load()["providers"])
def set(self, payload: ProviderIn) -> dict:
ensure_tables()
table = get_table(PROVIDERS_TABLE)
existing = table.find_one(name=payload.name)
record = {
"name": payload.name,
"base_url": payload.base_url,
"api_key": payload.api_key,
"is_active": 1 if payload.is_active else 0,
"updated_at": _now(),
}
if existing:
table.update({**record, "id": existing["id"]}, ["id"])
else:
record["created_at"] = _now()
table.insert(record)
bump_cache_version(CACHE_NAME)
_ROUTING_CACHE.clear()
return self.get(payload.name) or record
def remove(self, name: str) -> bool:
name = (name or "").strip().lower()
if not name or PROVIDERS_TABLE not in db.tables:
return False
removed = int(get_table(PROVIDERS_TABLE).delete(name=name))
if removed:
bump_cache_version(CACHE_NAME)
_ROUTING_CACHE.clear()
return bool(removed)
class ModelStore:
def list(self) -> list[dict]:
rows = []
for route in _load()["models"].values():
rows.append(route.__dict__.copy())
return sorted(rows, key=lambda r: r["source_model"])
def get(self, source_model: str) -> Optional[ModelRoute]:
if not source_model:
return None
return _load()["models"].get(source_model.strip())
def resolve(self, source_model: Optional[str], kind: str) -> Optional[ModelRoute]:
if not source_model:
return None
route = _load()["models"].get(source_model.strip())
if route is None or not route.is_active or route.kind != kind:
return None
return route
def count(self) -> int:
return len(_load()["models"])
def set(self, payload: ModelRouteIn) -> dict:
ensure_tables()
table = get_table(MODELS_TABLE)
existing = table.find_one(source_model=payload.source_model)
record = {
"source_model": payload.source_model,
"provider": payload.provider,
"target_model": payload.target_model,
"kind": payload.kind,
"vision_provider": payload.vision_provider,
"vision_model": payload.vision_model,
"context_window": payload.context_window,
"price_cache_hit_per_m": payload.price_cache_hit_per_m,
"price_cache_miss_per_m": payload.price_cache_miss_per_m,
"price_output_per_m": payload.price_output_per_m,
"price_input_per_m": payload.price_input_per_m,
"is_active": 1 if payload.is_active else 0,
"updated_at": _now(),
}
if existing:
table.update({**record, "id": existing["id"]}, ["id"])
else:
record["created_at"] = _now()
table.insert(record)
bump_cache_version(CACHE_NAME)
_ROUTING_CACHE.clear()
route = self.get(payload.source_model)
return route.__dict__.copy() if route else record
def remove(self, source_model: str) -> bool:
source_model = (source_model or "").strip()
if not source_model or MODELS_TABLE not in db.tables:
return False
removed = int(get_table(MODELS_TABLE).delete(source_model=source_model))
if removed:
bump_cache_version(CACHE_NAME)
_ROUTING_CACHE.clear()
return bool(removed)
provider_store = ProviderStore()
model_store = ModelStore()
def _provider_overlay(name: str, base_key: str, url_key: str, overlay: dict) -> None:
provider = provider_store.get(name)
if provider is None:
return
if provider.get("base_url"):
overlay[url_key] = provider["base_url"]
if provider.get("api_key"):
overlay[base_key] = provider["api_key"]
def chat_overlay(requested_model: Optional[str], base_cfg: dict) -> Optional[dict]:
route = model_store.resolve(requested_model, "chat")
if route is None:
return None
overlay: dict = {
"gateway_force_model": True,
"gateway_model": route.target_model,
"gateway_price_cache_hit_per_m": route.price_cache_hit_per_m,
"gateway_price_cache_miss_per_m": route.price_cache_miss_per_m,
"gateway_price_output_per_m": route.price_output_per_m,
}
if route.provider:
_provider_overlay(
route.provider, "gateway_api_key", "gateway_upstream_url", overlay
)
if route.context_window:
from devplacepy.services.openai_gateway.usage import parse_context_map
context_map = parse_context_map(base_cfg.get("gateway_model_context_map"))
context_map[route.target_model] = route.context_window
overlay["gateway_model_context_map"] = context_map
if route.vision_model:
overlay["gateway_vision_enabled"] = True
overlay["gateway_vision_model"] = route.vision_model
overlay["gateway_vision_price_input_per_m"] = route.price_input_per_m
overlay["gateway_vision_price_output_per_m"] = route.price_output_per_m
vision_provider = route.vision_provider or route.provider
if vision_provider:
_provider_overlay(
vision_provider, "gateway_vision_key", "gateway_vision_url", overlay
)
return overlay
def embed_overlay(requested_model: Optional[str], base_cfg: dict) -> Optional[dict]:
route = model_store.resolve(requested_model, "embed")
if route is None:
return None
overlay: dict = {
"gateway_force_model": True,
"gateway_embed_model": route.target_model,
"gateway_embed_price_input_per_m": route.price_input_per_m,
}
provider = provider_store.get(route.provider) if route.provider else None
if provider:
if provider.get("base_url"):
overlay["gateway_embed_url"] = _embed_url_from_base(provider["base_url"])
if provider.get("api_key"):
overlay["gateway_embed_key"] = provider["api_key"]
return overlay