488 lines
20 KiB
Python
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),
|
|
)
|