deerflow-code/offline-backend-20260512/backend/app/gateway/routers/model_management.py
2026-09-07 18:24:55 +08:00

488 lines
20 KiB
Python

"""Administrator model configuration backed by DeerFlow's ``config.yaml``.
The public ``/api/models`` endpoint intentionally stays small and never
returns secrets. This router is the complementary administrator surface: it
edits only the top-level ``models:`` section in the active DeerFlow config,
keeps a disk backup before every write, and swaps the gateway's runtime config
after validation so new requests can use the change without a restart.
"""
from __future__ import annotations
import asyncio
import copy
import logging
import os
import re
import shutil
import tempfile
from datetime import UTC, datetime
from pathlib import Path
from time import perf_counter
from typing import Any
import yaml
from fastapi import APIRouter, HTTPException, Request
from langchain_core.messages import HumanMessage
from pydantic import BaseModel, ConfigDict, Field, field_validator
from app.gateway.deps import get_optional_user_from_request
from deerflow.config.app_config import AppConfig, reload_app_config
from deerflow.config.model_config import ModelConfig
from deerflow.models.factory import create_chat_model
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/api/admin/models", tags=["admin-models"])
_MODEL_BLOCK_RE = re.compile(r"(?m)^models:[^\r\n]*(?:\r?\n|$)")
_TOP_LEVEL_KEY_RE = re.compile(r"(?m)^[A-Za-z_][A-Za-z0-9_-]*:[^\r\n]*(?:\r?\n|$)")
_BACKUP_DIR_NAME = ".model-config-backups"
_MUTATION_LOCK = asyncio.Lock()
_EDITOR_FIELDS = frozenset(
{
"name",
"display_name",
"description",
"use",
"model",
"api_key",
"base_url",
"request_timeout",
"max_retries",
"max_tokens",
"temperature",
"stream_usage",
"supports_thinking",
"supports_reasoning_effort",
"supports_vision",
"when_thinking_enabled",
"when_thinking_disabled",
"thinking",
"use_responses_api",
"output_version",
}
)
_PRESERVED_FIELDS = frozenset(ModelConfig.model_fields) | _EDITOR_FIELDS
class ModelAdminItem(BaseModel):
"""A model configuration safe to return to an administrator browser."""
name: str
display_name: str | None = None
description: str | None = None
use: str
model: str
base_url: str | None = None
request_timeout: float | None = None
max_retries: int | None = None
max_tokens: int | None = None
temperature: float | None = None
stream_usage: bool | None = None
supports_thinking: bool = False
supports_reasoning_effort: bool = False
supports_vision: bool = False
when_thinking_enabled: dict[str, Any] | None = None
when_thinking_disabled: dict[str, Any] | None = None
thinking: dict[str, Any] | None = None
use_responses_api: bool | None = None
output_version: str | None = None
extra_settings: dict[str, Any] = Field(default_factory=dict)
api_key_configured: bool = False
api_key_source: str | None = None
source: str = "config.yaml"
class ModelAdminListResponse(BaseModel):
models: list[ModelAdminItem]
config_path: str
updated_at: str | None = None
class ModelUpsertRequest(BaseModel):
"""Editable subset of the ``ModelConfig`` schema plus safe extra fields."""
model_config = ConfigDict(extra="forbid")
name: str = Field(min_length=1, max_length=160)
display_name: str | None = Field(default=None, max_length=160)
description: str | None = Field(default=None, max_length=1000)
use: str = Field(min_length=1, max_length=255)
model: str = Field(min_length=1, max_length=255)
# ``None`` means "keep the currently configured secret" on update.
# The value is write-only in practice: it is never part of a response.
api_key: str | None = Field(default=None, max_length=4096)
clear_api_key: bool = False
base_url: str | None = Field(default=None, max_length=2048)
request_timeout: float | None = Field(default=None, ge=1, le=600)
max_retries: int | None = Field(default=None, ge=0, le=20)
max_tokens: int | None = Field(default=None, ge=1, le=1_000_000)
temperature: float | None = Field(default=None, ge=0, le=2)
stream_usage: bool | None = None
supports_thinking: bool = False
supports_reasoning_effort: bool = False
supports_vision: bool = False
when_thinking_enabled: dict[str, Any] | None = None
when_thinking_disabled: dict[str, Any] | None = None
thinking: dict[str, Any] | None = None
use_responses_api: bool | None = None
output_version: str | None = Field(default=None, max_length=128)
extra_settings: dict[str, Any] = Field(default_factory=dict)
@field_validator("name", "use", "model")
@classmethod
def _trim_required(cls, value: str) -> str:
normalized = value.strip()
if not normalized:
raise ValueError("value cannot be empty")
return normalized
@field_validator("display_name", "description", "base_url", "api_key", "output_version")
@classmethod
def _trim_optional(cls, value: str | None) -> str | None:
return value.strip() if isinstance(value, str) else value
class ModelTestRequest(BaseModel):
model: ModelUpsertRequest
class ModelOrderRequest(BaseModel):
"""The complete ordered set of model stable names."""
names: list[str] = Field(min_length=1, max_length=500)
class ModelTestResponse(BaseModel):
success: bool
latency_ms: int
message: str
preview: str | None = None
def _mask_env_source(value: Any) -> str | None:
if isinstance(value, str) and value.startswith("$") and len(value) > 1:
return value[1:]
return None
def _to_admin_item(raw: dict[str, Any]) -> ModelAdminItem:
api_key = raw.get("api_key")
extra = {
key: copy.deepcopy(value)
for key, value in raw.items()
if key not in _PRESERVED_FIELDS
}
return ModelAdminItem(
name=str(raw.get("name") or ""),
display_name=raw.get("display_name"),
description=raw.get("description"),
use=str(raw.get("use") or ""),
model=str(raw.get("model") or ""),
base_url=raw.get("base_url"),
request_timeout=raw.get("request_timeout"),
max_retries=raw.get("max_retries"),
max_tokens=raw.get("max_tokens"),
temperature=raw.get("temperature"),
stream_usage=raw.get("stream_usage"),
supports_thinking=bool(raw.get("supports_thinking", False)),
supports_reasoning_effort=bool(raw.get("supports_reasoning_effort", False)),
supports_vision=bool(raw.get("supports_vision", False)),
when_thinking_enabled=copy.deepcopy(raw.get("when_thinking_enabled")),
when_thinking_disabled=copy.deepcopy(raw.get("when_thinking_disabled")),
thinking=copy.deepcopy(raw.get("thinking")),
use_responses_api=raw.get("use_responses_api"),
output_version=raw.get("output_version"),
extra_settings=extra,
api_key_configured=bool(api_key),
api_key_source=_mask_env_source(api_key),
)
def _resolve_config_path() -> Path:
return AppConfig.resolve_config_path()
def _read_config_data(path: Path) -> tuple[str, dict[str, Any], list[dict[str, Any]]]:
raw_text = path.read_text(encoding="utf-8")
try:
data = yaml.safe_load(raw_text) or {}
except yaml.YAMLError as exc:
raise HTTPException(status_code=500, detail=f"Unable to parse active config.yaml: {exc}") from exc
if not isinstance(data, dict):
raise HTTPException(status_code=500, detail="The active config.yaml root must be a mapping")
raw_models = data.get("models") or []
if not isinstance(raw_models, list) or not all(isinstance(item, dict) for item in raw_models):
raise HTTPException(status_code=500, detail="config.yaml models must be a list of mappings")
return raw_text, data, [copy.deepcopy(item) for item in raw_models]
def _render_models_block(models: list[dict[str, Any]], newline: str) -> str:
rendered = yaml.safe_dump(
{"models": models},
allow_unicode=True,
sort_keys=False,
default_flow_style=False,
width=120,
)
if newline != "\n":
rendered = rendered.replace("\n", newline)
return rendered if rendered.endswith(newline) else f"{rendered}{newline}"
def _replace_models_block(raw_text: str, models: list[dict[str, Any]]) -> str:
"""Replace the active model entries without touching other config sections.
PyYAML does not retain comments. Keep every comment found in the old
model section as documentation after the newly rendered model list rather
than silently deleting an operator's notes or disabled examples.
"""
block_match = _MODEL_BLOCK_RE.search(raw_text)
if block_match is None:
raise HTTPException(status_code=500, detail="The active config.yaml has no top-level models section")
newline = "\r\n" if "\r\n" in raw_text else "\n"
next_key = _TOP_LEVEL_KEY_RE.search(raw_text, block_match.end())
end = next_key.start() if next_key else len(raw_text)
previous_section = raw_text[block_match.end() : end]
preserved_comments = "".join(
line
for line in previous_section.splitlines(keepends=True)
if line.lstrip().startswith("#")
)
rendered = _render_models_block(models, newline)
if preserved_comments:
if not preserved_comments.endswith(("\n", "\r")):
preserved_comments = f"{preserved_comments}{newline}"
rendered = f"{rendered}{preserved_comments}"
return f"{raw_text[:block_match.start()]}{rendered}{raw_text[end:]}"
def _backup_and_write_config(path: Path, raw_before: str, raw_after: str) -> Path:
backup_dir = path.parent / _BACKUP_DIR_NAME
backup_dir.mkdir(parents=True, exist_ok=True)
stamp = datetime.now(UTC).strftime("%Y%m%dT%H%M%S%fZ")
backup_path = backup_dir / f"{path.stem}.{stamp}.yaml"
# Persist the exact text that was read for this mutation. This makes the
# rollback deterministic even if another operator changes the file while
# the gateway is processing a request.
with backup_path.open("w", encoding="utf-8", newline="") as handle:
handle.write(raw_before)
fd, temp_name = tempfile.mkstemp(prefix=f".{path.name}.", suffix=".tmp", dir=path.parent)
try:
with os.fdopen(fd, "w", encoding="utf-8", newline="") as handle:
handle.write(raw_after)
handle.flush()
os.fsync(handle.fileno())
os.replace(temp_name, path)
except Exception:
try:
os.unlink(temp_name)
except OSError:
pass
raise
return backup_path
def _restore_backup(path: Path, backup_path: Path) -> None:
try:
shutil.copy2(backup_path, path)
except Exception:
logger.exception("Failed to restore model config backup at %s", backup_path)
def _request_to_raw_model(payload: ModelUpsertRequest, existing: dict[str, Any] | None = None) -> dict[str, Any]:
merged = copy.deepcopy(existing or {})
supplied = payload.model_dump(exclude={"extra_settings", "clear_api_key"}, exclude_none=True)
merged.update(supplied)
for key, value in payload.extra_settings.items():
normalized = str(key).strip()
if not normalized:
raise HTTPException(status_code=422, detail="Extra setting names cannot be blank")
if normalized in _EDITOR_FIELDS or normalized in {"api_key", "clear_api_key", "extra_settings"}:
raise HTTPException(status_code=422, detail=f"Extra setting '{normalized}' duplicates a managed field")
merged[normalized] = copy.deepcopy(value)
if payload.clear_api_key:
merged.pop("api_key", None)
elif payload.api_key is None and existing is not None and "api_key" in existing:
merged["api_key"] = existing["api_key"]
# Validate the raw schema without resolving `$ENV_VAR` values. Environment
# resolution happens during the post-write runtime reload and is rolled back
# if unavailable.
try:
ModelConfig.model_validate(merged)
except Exception as exc:
raise HTTPException(status_code=422, detail=f"Invalid model configuration: {exc}") from exc
return merged
def _redact_text(value: str, secrets: list[str]) -> str:
result = value
for secret in secrets:
if secret:
result = result.replace(secret, "***")
return result[:800]
async def _require_admin(request: Request) -> None:
user = getattr(request.state, "user", None)
if user is None:
user = await get_optional_user_from_request(request)
if user is None:
raise HTTPException(status_code=401, detail="Authentication is required")
if getattr(user, "system_role", None) != "admin":
raise HTTPException(status_code=403, detail="Model configuration requires an administrator")
async def _save_models(request: Request, raw_before: str, models: list[dict[str, Any]]) -> list[dict[str, Any]]:
path = _resolve_config_path()
raw_after = _replace_models_block(raw_before, models)
backup_path = await asyncio.to_thread(_backup_and_write_config, path, raw_before, raw_after)
try:
config = reload_app_config(str(path))
request.app.state.config = config
except Exception as exc:
logger.warning("Updated model configuration could not be activated: %s", type(exc).__name__)
await asyncio.to_thread(_restore_backup, path, backup_path)
try:
request.app.state.config = reload_app_config(str(path))
except Exception:
logger.exception("Failed to reload restored model config")
raise HTTPException(
status_code=422,
detail="The model configuration could not be activated; the previous config was restored.",
) from exc
return models
@router.get("", response_model=ModelAdminListResponse)
async def list_model_configurations(request: Request) -> ModelAdminListResponse:
await _require_admin(request)
path = _resolve_config_path()
_, _, models = await asyncio.to_thread(_read_config_data, path)
updated_at = datetime.fromtimestamp(path.stat().st_mtime, tz=UTC).isoformat()
return ModelAdminListResponse(
models=[_to_admin_item(item) for item in models],
config_path=str(path),
updated_at=updated_at,
)
@router.post("", response_model=ModelAdminItem, status_code=201)
async def create_model_configuration(body: ModelUpsertRequest, request: Request) -> ModelAdminItem:
await _require_admin(request)
async with _MUTATION_LOCK:
path = _resolve_config_path()
raw_before, _, models = await asyncio.to_thread(_read_config_data, path)
if any(str(item.get("name")) == body.name for item in models):
raise HTTPException(status_code=409, detail=f"A model named '{body.name}' already exists")
created = _request_to_raw_model(body)
await _save_models(request, raw_before, [*models, created])
return _to_admin_item(created)
@router.put("/order", response_model=list[ModelAdminItem])
async def reorder_model_configurations(body: ModelOrderRequest, request: Request) -> list[ModelAdminItem]:
"""Persist the configured model order.
The first item is intentionally the deployment default: public model
selectors preserve the ``config.yaml`` order and use their first option
when no model was explicitly chosen. Requiring a complete permutation
prevents a stale browser from accidentally dropping a newly added model.
"""
await _require_admin(request)
if len(body.names) != len(set(body.names)):
raise HTTPException(status_code=422, detail="Each model may appear only once in the order")
async with _MUTATION_LOCK:
path = _resolve_config_path()
raw_before, _, models = await asyncio.to_thread(_read_config_data, path)
configured_names = [str(item.get("name") or "") for item in models]
if len(body.names) != len(configured_names) or set(body.names) != set(configured_names):
raise HTTPException(
status_code=409,
detail="The model list changed. Refresh the page and try again.",
)
models_by_name = {str(item.get("name")): item for item in models}
ordered_models = [models_by_name[name] for name in body.names]
await _save_models(request, raw_before, ordered_models)
return [_to_admin_item(item) for item in ordered_models]
@router.put("/{model_name:path}", response_model=ModelAdminItem)
async def update_model_configuration(model_name: str, body: ModelUpsertRequest, request: Request) -> ModelAdminItem:
await _require_admin(request)
if body.name != model_name:
raise HTTPException(status_code=422, detail="A model's stable name cannot be changed; edit its display name instead")
async with _MUTATION_LOCK:
path = _resolve_config_path()
raw_before, _, models = await asyncio.to_thread(_read_config_data, path)
index = next((i for i, item in enumerate(models) if str(item.get("name")) == model_name), None)
if index is None:
raise HTTPException(status_code=404, detail=f"Model '{model_name}' was not found")
updated = _request_to_raw_model(body, models[index])
next_models = [*models]
next_models[index] = updated
await _save_models(request, raw_before, next_models)
return _to_admin_item(updated)
@router.delete("/{model_name:path}", status_code=204)
async def delete_model_configuration(model_name: str, request: Request) -> None:
"""Delete one model configuration and activate the remaining model list."""
await _require_admin(request)
async with _MUTATION_LOCK:
path = _resolve_config_path()
raw_before, _, models = await asyncio.to_thread(_read_config_data, path)
index = next((i for i, item in enumerate(models) if str(item.get("name")) == model_name), None)
if index is None:
raise HTTPException(status_code=404, detail=f"Model '{model_name}' was not found")
await _save_models(request, raw_before, [*models[:index], *models[index + 1 :]])
@router.post("/test", response_model=ModelTestResponse)
async def test_model_connection(body: ModelTestRequest, request: Request) -> ModelTestResponse:
"""Perform one bounded minimal completion without persisting the draft."""
await _require_admin(request)
path = _resolve_config_path()
_, _, models = await asyncio.to_thread(_read_config_data, path)
existing = next((item for item in models if str(item.get("name")) == body.model.name), None)
candidate = _request_to_raw_model(body.model, existing)
candidate_secret = str(candidate.get("api_key") or "")
secrets = [candidate_secret]
try:
runtime_candidate = AppConfig.resolve_env_variables(copy.deepcopy(candidate))
# An environment-variable API key is resolved only in memory for this
# request. Redact both forms because provider exceptions can include
# either the original `$NAME` reference or its resolved value.
resolved_secret = str(runtime_candidate.get("api_key") or "")
secrets = [candidate_secret, resolved_secret]
model_config = ModelConfig.model_validate(runtime_candidate)
base_config = request.app.state.config
test_config = base_config.model_copy(update={"models": [model_config]})
chat_model = create_chat_model(model_config.name, app_config=test_config, force_disable_thinking=True)
started = perf_counter()
response = await asyncio.wait_for(
chat_model.ainvoke([HumanMessage(content="Reply with exactly: OK")]),
timeout=min(float(candidate.get("request_timeout") or 30), 45),
)
latency_ms = round((perf_counter() - started) * 1000)
preview = _redact_text(str(getattr(response, "content", "OK")), secrets)
return ModelTestResponse(success=True, latency_ms=latency_ms, message="Connection successful", preview=preview)
except Exception as exc:
# Provider errors sometimes echo a credential. Never log the full
# exception or return it without redacting the submitted key.
logger.warning("Model connection test failed for '%s': %s", candidate.get("name"), type(exc).__name__)
return ModelTestResponse(
success=False,
latency_ms=0,
message=_redact_text(str(exc) or type(exc).__name__, secrets),
)