forked from retoor/devplacepy
Admin-unlimited Dev Workspaces: an admin-owned workspace is now exempt from the max-workspace-count limit, the max-tunnel-count limit, and the whole idle-stop/idle-warn/retention-delete lifecycle. Resolved once in quota.resolve() as Limits.unlimited (owner uid checked against get_admin_uids()), consumed at the three enforcement points (provision.ensure, provision.publish_tunnel, WorkspaceService._advance_lifecycle). Also hardens get_admin_uids()/get_primary_admin_uid() against a partially-schemaed users table (uid/role column guard), which a fresh test/init_db() path could hit. AI gateway per-model automatic fallback: any gateway_models route (chat/embed/image) can now name a fallback_model, picked on /admin/gateway from a select box of other configured public model names of the same kind only (never an internal upstream model id). When a route fails after its own retries are exhausted, the gateway retries once, automatically, against the fallback's own provider/pricing/key, before any bytes reach the client (including for a streaming response). One hop only, no chains or cycles; self-reference and cross-kind fallbacks are rejected at write time. AI gateway real upstream streaming and thinking-default control: stream:true is now forwarded to the upstream and relayed to the client as real SSE chunks (measured TTFT/inter-token latency) instead of a simulated split response, and every chat/vision call explicitly disables model "thinking" by default (admin-overridable via gateway_thinking), with per-dialect handling for DeepSeek, OpenRouter, and Ollama. Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01TjdKTWgWpW2SMNW8SFqxz5
707 lines
26 KiB
Python
707 lines
26 KiB
Python
# 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, model_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", "image")
|
|
|
|
_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 _image_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")] + "/images"
|
|
return base_url
|
|
|
|
|
|
MODEL_TIER2_COLUMNS = (
|
|
"context_tier_threshold_tokens",
|
|
"price_cache_hit_per_m_tier2",
|
|
"price_cache_miss_per_m_tier2",
|
|
"price_output_per_m_tier2",
|
|
"price_input_per_m_tier2",
|
|
"off_peak_start_minute",
|
|
"off_peak_end_minute",
|
|
"off_peak_discount_pct",
|
|
)
|
|
|
|
|
|
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, "
|
|
"context_tier_threshold_tokens INTEGER DEFAULT 0, "
|
|
"price_cache_hit_per_m_tier2 REAL, price_cache_miss_per_m_tier2 REAL, "
|
|
"price_output_per_m_tier2 REAL, price_input_per_m_tier2 REAL, "
|
|
"off_peak_start_minute INTEGER, off_peak_end_minute INTEGER, "
|
|
"off_peak_discount_pct REAL DEFAULT 0, fallback_model TEXT DEFAULT '', "
|
|
"is_active INTEGER DEFAULT 1, created_at TEXT, updated_at TEXT)"
|
|
)
|
|
models_table = get_table(MODELS_TABLE)
|
|
for column in MODEL_TIER2_COLUMNS:
|
|
if not models_table.has_column(column):
|
|
models_table.create_column_by_example(column, 0.0)
|
|
if not models_table.has_column("fallback_model"):
|
|
models_table.create_column_by_example("fallback_model", "")
|
|
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)
|
|
context_tier_threshold_tokens: int = Field(default=0, ge=0, le=100_000_000)
|
|
price_cache_hit_per_m_tier2: Optional[float] = Field(default=None, ge=0)
|
|
price_cache_miss_per_m_tier2: Optional[float] = Field(default=None, ge=0)
|
|
price_output_per_m_tier2: Optional[float] = Field(default=None, ge=0)
|
|
price_input_per_m_tier2: Optional[float] = Field(default=None, ge=0)
|
|
off_peak_start_minute: Optional[int] = Field(default=None, ge=0, le=1439)
|
|
off_peak_end_minute: Optional[int] = Field(default=None, ge=0, le=1439)
|
|
off_peak_discount_pct: float = Field(default=0.0, ge=0, le=100)
|
|
fallback_model: str = Field(default="", max_length=128)
|
|
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", "fallback_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', 'embed', or 'image'")
|
|
return value
|
|
|
|
@model_validator(mode="after")
|
|
def _check_off_peak_window(self) -> "ModelRouteIn":
|
|
has_start = self.off_peak_start_minute is not None
|
|
has_end = self.off_peak_end_minute is not None
|
|
if has_start != has_end:
|
|
raise ValueError(
|
|
"Off-peak start and end minute must both be set, or both left blank"
|
|
)
|
|
return self
|
|
|
|
@model_validator(mode="after")
|
|
def _check_fallback_is_not_self(self) -> "ModelRouteIn":
|
|
if self.fallback_model and self.fallback_model == self.source_model:
|
|
raise ValueError("A model cannot fall back to itself")
|
|
return self
|
|
|
|
|
|
@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
|
|
context_tier_threshold_tokens: int
|
|
price_cache_hit_per_m_tier2: Optional[float]
|
|
price_cache_miss_per_m_tier2: Optional[float]
|
|
price_output_per_m_tier2: Optional[float]
|
|
price_input_per_m_tier2: Optional[float]
|
|
off_peak_start_minute: Optional[int]
|
|
off_peak_end_minute: Optional[int]
|
|
off_peak_discount_pct: float
|
|
fallback_model: str
|
|
is_active: bool
|
|
|
|
|
|
def _opt_float(row: dict, key: str) -> Optional[float]:
|
|
value = row.get(key)
|
|
return float(value) if value is not None else None
|
|
|
|
|
|
def _opt_int(row: dict, key: str) -> Optional[int]:
|
|
value = row.get(key)
|
|
return int(value) if value is not None else None
|
|
|
|
|
|
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),
|
|
context_tier_threshold_tokens=int(
|
|
row.get("context_tier_threshold_tokens") or 0
|
|
),
|
|
price_cache_hit_per_m_tier2=_opt_float(row, "price_cache_hit_per_m_tier2"),
|
|
price_cache_miss_per_m_tier2=_opt_float(row, "price_cache_miss_per_m_tier2"),
|
|
price_output_per_m_tier2=_opt_float(row, "price_output_per_m_tier2"),
|
|
price_input_per_m_tier2=_opt_float(row, "price_input_per_m_tier2"),
|
|
off_peak_start_minute=_opt_int(row, "off_peak_start_minute"),
|
|
off_peak_end_minute=_opt_int(row, "off_peak_end_minute"),
|
|
off_peak_discount_pct=float(row.get("off_peak_discount_pct") or 0.0),
|
|
fallback_model=str(row.get("fallback_model") or ""),
|
|
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,
|
|
"context_tier_threshold_tokens": payload.context_tier_threshold_tokens,
|
|
"price_cache_hit_per_m_tier2": payload.price_cache_hit_per_m_tier2,
|
|
"price_cache_miss_per_m_tier2": payload.price_cache_miss_per_m_tier2,
|
|
"price_output_per_m_tier2": payload.price_output_per_m_tier2,
|
|
"price_input_per_m_tier2": payload.price_input_per_m_tier2,
|
|
"off_peak_start_minute": payload.off_peak_start_minute,
|
|
"off_peak_end_minute": payload.off_peak_end_minute,
|
|
"off_peak_discount_pct": payload.off_peak_discount_pct,
|
|
"fallback_model": payload.fallback_model,
|
|
"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 resolve_fallback(source_model: Optional[str], kind: str) -> Optional[ModelRoute]:
|
|
route = model_store.resolve(source_model, kind)
|
|
if route is None or not route.fallback_model:
|
|
return None
|
|
if route.fallback_model == route.source_model:
|
|
return None
|
|
fallback = model_store.resolve(route.fallback_model, kind)
|
|
if fallback is None or fallback.source_model == route.source_model:
|
|
return None
|
|
return fallback
|
|
|
|
DEEPSEEK_DEFAULT_ROUTES = (
|
|
{
|
|
"source_model": "deepseek-v4-flash",
|
|
"target_model": "deepseek-v4-flash",
|
|
"context_window": 1_048_576,
|
|
"price_cache_hit_per_m": 0.0028,
|
|
"price_cache_miss_per_m": 0.14,
|
|
"price_output_per_m": 0.28,
|
|
},
|
|
{
|
|
"source_model": "deepseek-v4-pro",
|
|
"target_model": "deepseek-v4-pro",
|
|
"context_window": 1_048_576,
|
|
"price_cache_hit_per_m": 0.003625,
|
|
"price_cache_miss_per_m": 0.435,
|
|
"price_output_per_m": 0.87,
|
|
},
|
|
{
|
|
"source_model": "molodetz",
|
|
"target_model": "deepseek-v4-flash",
|
|
"context_window": 1_048_576,
|
|
"price_cache_hit_per_m": 0.0028,
|
|
"price_cache_miss_per_m": 0.14,
|
|
"price_output_per_m": 0.28,
|
|
},
|
|
{
|
|
"source_model": "molodetz-pro",
|
|
"target_model": "deepseek-v4-pro",
|
|
"context_window": 1_048_576,
|
|
"price_cache_hit_per_m": 0.003625,
|
|
"price_cache_miss_per_m": 0.435,
|
|
"price_output_per_m": 0.87,
|
|
},
|
|
)
|
|
|
|
|
|
def seed_default_deepseek_routes() -> None:
|
|
ensure_tables()
|
|
table = get_table(MODELS_TABLE)
|
|
seeded = False
|
|
for defaults in DEEPSEEK_DEFAULT_ROUTES:
|
|
if table.find_one(source_model=defaults["source_model"]):
|
|
continue
|
|
record = {
|
|
"source_model": defaults["source_model"],
|
|
"provider": "",
|
|
"target_model": defaults["target_model"],
|
|
"kind": "chat",
|
|
"vision_provider": "",
|
|
"vision_model": "",
|
|
"context_window": defaults["context_window"],
|
|
"price_cache_hit_per_m": defaults["price_cache_hit_per_m"],
|
|
"price_cache_miss_per_m": defaults["price_cache_miss_per_m"],
|
|
"price_output_per_m": defaults["price_output_per_m"],
|
|
"price_input_per_m": 0.0,
|
|
"context_tier_threshold_tokens": 0,
|
|
"price_cache_hit_per_m_tier2": None,
|
|
"price_cache_miss_per_m_tier2": None,
|
|
"price_output_per_m_tier2": None,
|
|
"price_input_per_m_tier2": None,
|
|
"off_peak_start_minute": None,
|
|
"off_peak_end_minute": None,
|
|
"off_peak_discount_pct": 0.0,
|
|
"is_active": 1,
|
|
"created_at": _now(),
|
|
"updated_at": _now(),
|
|
}
|
|
table.insert(record)
|
|
seeded = True
|
|
if seeded:
|
|
bump_cache_version(CACHE_NAME)
|
|
_ROUTING_CACHE.clear()
|
|
|
|
|
|
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,
|
|
"gateway_price_cache_hit_per_m_tier2": route.price_cache_hit_per_m_tier2,
|
|
"gateway_price_cache_miss_per_m_tier2": route.price_cache_miss_per_m_tier2,
|
|
"gateway_price_output_per_m_tier2": route.price_output_per_m_tier2,
|
|
"gateway_context_tier_threshold_tokens": route.context_tier_threshold_tokens,
|
|
"gateway_off_peak_start_minute": route.off_peak_start_minute,
|
|
"gateway_off_peak_end_minute": route.off_peak_end_minute,
|
|
"gateway_off_peak_discount_pct": route.off_peak_discount_pct,
|
|
}
|
|
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
|
|
overlay["gateway_vision_price_input_per_m_tier2"] = (
|
|
route.price_input_per_m_tier2
|
|
)
|
|
overlay["gateway_vision_price_output_per_m_tier2"] = (
|
|
route.price_output_per_m_tier2
|
|
)
|
|
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,
|
|
"gateway_embed_price_input_per_m_tier2": route.price_input_per_m_tier2,
|
|
"gateway_context_tier_threshold_tokens": route.context_tier_threshold_tokens,
|
|
"gateway_off_peak_start_minute": route.off_peak_start_minute,
|
|
"gateway_off_peak_end_minute": route.off_peak_end_minute,
|
|
"gateway_off_peak_discount_pct": route.off_peak_discount_pct,
|
|
}
|
|
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
|
|
|
|
|
|
def image_overlay(requested_model: Optional[str], base_cfg: dict) -> Optional[dict]:
|
|
route = model_store.resolve(requested_model, "image")
|
|
if route is None:
|
|
return None
|
|
overlay: dict = {
|
|
"gateway_force_model": True,
|
|
"gateway_image_model": route.target_model,
|
|
"gateway_image_price_per_call": route.price_input_per_m,
|
|
"gateway_off_peak_start_minute": route.off_peak_start_minute,
|
|
"gateway_off_peak_end_minute": route.off_peak_end_minute,
|
|
"gateway_off_peak_discount_pct": route.off_peak_discount_pct,
|
|
}
|
|
provider = provider_store.get(route.provider) if route.provider else None
|
|
if provider:
|
|
if provider.get("base_url"):
|
|
overlay["gateway_image_url"] = _image_url_from_base(provider["base_url"])
|
|
if provider.get("api_key"):
|
|
overlay["gateway_image_key"] = provider["api_key"]
|
|
return overlay
|
|
|
|
|
|
OPENROUTER_PROVIDER_NAME = "openrouter"
|
|
FLUX_TARGET_MODEL = "black-forest-labs/flux.2-pro"
|
|
IMAGE_SOURCE_MODEL = "molodetz-img-small"
|
|
RETIRED_IMAGE_MODELS = {
|
|
"black-forest-labs/flux-1.1-pro": FLUX_TARGET_MODEL,
|
|
"black-forest-labs/flux-schnell": "black-forest-labs/flux.2-klein-4b",
|
|
}
|
|
LEGACY_IMAGE_URL = "https://openrouter.ai/api/v1/images/generations"
|
|
|
|
|
|
def seed_default_image_routes() -> None:
|
|
import os
|
|
|
|
from devplacepy.services.openai_gateway import config as gw_config
|
|
|
|
ensure_tables()
|
|
openrouter_key = os.environ.get("OPENROUTER_API_KEY", "")
|
|
if openrouter_key:
|
|
existing_provider = provider_store.get(OPENROUTER_PROVIDER_NAME)
|
|
if existing_provider is None:
|
|
provider_store.set(
|
|
ProviderIn(
|
|
name=OPENROUTER_PROVIDER_NAME,
|
|
base_url="https://openrouter.ai/api/v1/chat/completions",
|
|
api_key=openrouter_key,
|
|
is_active=True,
|
|
)
|
|
)
|
|
models_table = get_table(MODELS_TABLE)
|
|
if models_table.find_one(source_model=IMAGE_SOURCE_MODEL):
|
|
return
|
|
record = {
|
|
"source_model": IMAGE_SOURCE_MODEL,
|
|
"provider": OPENROUTER_PROVIDER_NAME if openrouter_key else "",
|
|
"target_model": FLUX_TARGET_MODEL,
|
|
"kind": "image",
|
|
"vision_provider": "",
|
|
"vision_model": "",
|
|
"context_window": 0,
|
|
"price_cache_hit_per_m": 0.0,
|
|
"price_cache_miss_per_m": 0.0,
|
|
"price_output_per_m": 0.0,
|
|
"price_input_per_m": gw_config.IMAGE_PRICE_PER_CALL_DEFAULT,
|
|
"context_tier_threshold_tokens": 0,
|
|
"price_cache_hit_per_m_tier2": None,
|
|
"price_cache_miss_per_m_tier2": None,
|
|
"price_output_per_m_tier2": None,
|
|
"price_input_per_m_tier2": None,
|
|
"off_peak_start_minute": None,
|
|
"off_peak_end_minute": None,
|
|
"off_peak_discount_pct": 0.0,
|
|
"is_active": 1,
|
|
"created_at": _now(),
|
|
"updated_at": _now(),
|
|
}
|
|
models_table.insert(record)
|
|
bump_cache_version(CACHE_NAME)
|
|
_ROUTING_CACHE.clear()
|
|
|
|
|
|
def migrate_retired_image_gateway() -> None:
|
|
from devplacepy.database import get_setting, set_setting
|
|
from devplacepy.services.openai_gateway import config as gw_config
|
|
|
|
ensure_tables()
|
|
models_table = get_table(MODELS_TABLE)
|
|
now = _now()
|
|
changed = False
|
|
for row in models_table.find(kind="image"):
|
|
target = str(row.get("target_model") or "")
|
|
replacement = RETIRED_IMAGE_MODELS.get(target)
|
|
if not replacement:
|
|
continue
|
|
models_table.update(
|
|
{
|
|
"source_model": row["source_model"],
|
|
"target_model": replacement,
|
|
"updated_at": now,
|
|
},
|
|
["source_model"],
|
|
)
|
|
changed = True
|
|
image_url = get_setting("gateway_image_url", "")
|
|
if image_url in (LEGACY_IMAGE_URL, ""):
|
|
set_setting("gateway_image_url", gw_config.IMAGE_URL_DEFAULT)
|
|
changed = True
|
|
elif image_url.endswith("/images/generations"):
|
|
set_setting(
|
|
"gateway_image_url",
|
|
image_url[: -len("/generations")],
|
|
)
|
|
changed = True
|
|
image_model = get_setting("gateway_image_model", "")
|
|
replacement = RETIRED_IMAGE_MODELS.get(image_model)
|
|
if replacement:
|
|
set_setting("gateway_image_model", replacement)
|
|
changed = True
|
|
elif not image_model:
|
|
set_setting("gateway_image_model", gw_config.IMAGE_MODEL_DEFAULT)
|
|
changed = True
|
|
if changed:
|
|
bump_cache_version(CACHE_NAME)
|
|
_ROUTING_CACHE.clear()
|