forked from retoor/devplacepy
feat: add provider and model routing tables, admin UI, and audit category for gateway
Implement the multi-provider routing system for the OpenAI gateway, including two new database tables (`gateway_providers`, `gateway_models`) ensured at init, a new admin page at `/admin/gateway` with full CRUD for providers and model routes, and a `"gateway"` audit category mapped to `"ai"`. The routing layer sits transparently on top of the existing single-provider default path: unmatched model names fall through unchanged, while matched routes forward to the configured provider with their own pricing economy, vision model, and context window. Cross-worker cache invalidation uses a shared `_ROUTING_CACHE` bumped via `"gateway_routing"` cache version.
This commit is contained in:
@@ -13,6 +13,7 @@ from fastapi.responses import JSONResponse, Response, StreamingResponse
|
||||
from devplacepy import stealth
|
||||
from devplacepy.services.openai_gateway import config
|
||||
from devplacepy.services.openai_gateway.reliability import CircuitBreaker, retry_send
|
||||
from devplacepy.services.openai_gateway.routing import chat_overlay, embed_overlay
|
||||
from devplacepy.services.openai_gateway.system_message import apply_system_directives
|
||||
from devplacepy.services.openai_gateway.usage import (
|
||||
GatewayUsageLedger,
|
||||
@@ -213,6 +214,10 @@ class GatewayRuntime:
|
||||
self, body: dict, cfg: dict, owner: tuple, user_agent: str, log=None
|
||||
):
|
||||
log = log or (lambda message: None)
|
||||
overlay = chat_overlay(body.get("model"), cfg)
|
||||
if overlay:
|
||||
cfg = {**cfg, **overlay}
|
||||
log(f"routed model {body.get('model')!r} -> {cfg['gateway_model']!r}")
|
||||
client, sem = self._ensure(cfg)
|
||||
pricing = pricing_from_cfg(cfg)
|
||||
context_map = parse_context_map(cfg.get("gateway_model_context_map"))
|
||||
@@ -365,6 +370,10 @@ class GatewayRuntime:
|
||||
self, body: dict, cfg: dict, owner: tuple, user_agent: str, log=None
|
||||
):
|
||||
log = log or (lambda message: None)
|
||||
overlay = embed_overlay(body.get("model"), cfg)
|
||||
if overlay:
|
||||
cfg = {**cfg, **overlay}
|
||||
log(f"routed embed model {body.get('model')!r} -> {cfg['gateway_embed_model']!r}")
|
||||
if not cfg["gateway_embed_enabled"]:
|
||||
from devplacepy.services.audit import record as audit
|
||||
from devplacepy.services.openai_gateway.usage import audit_actor_for
|
||||
|
||||
@@ -0,0 +1,380 @@
|
||||
# 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
|
||||
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
|
||||
Reference in New Issue
Block a user