|
# 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
|