181 lines
7.8 KiB
Python
181 lines
7.8 KiB
Python
"""模型档位门控 + media variant 解析(从 app.py 析出,2026-07-23 拆分)。
|
|
|
|
web 层的「用户能不能用这个模型」判定收口:plan/role 查询、显式选择点的 403 门控、
|
|
image/video variant 的清单列举与解析。核心档位规则在 core/model_access.py,
|
|
这里只是 HTTP 语义包装(HTTPException)+ config 扫描。
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
from typing import Optional
|
|
from uuid import UUID
|
|
|
|
from fastapi import HTTPException
|
|
from sqlalchemy import select
|
|
|
|
from core.storage import session_scope
|
|
from core.storage.models import User
|
|
|
|
# 档位降级目标:存量 task 的模型已不在用户档位内时,下次起 run 落回这个(基线必含)。
|
|
FALLBACK_MODEL_PROFILE = "deepseek_v4.flash"
|
|
|
|
|
|
def user_plan_role(user_id: UUID) -> tuple[str, str]:
|
|
"""取该用户的 (plan, role);行不存在 → ("", "user")。模型访问门控用。"""
|
|
with session_scope() as s:
|
|
row = s.execute(
|
|
select(User.plan, User.role).where(User.user_id == user_id)
|
|
).first()
|
|
if row is None:
|
|
return "", "user"
|
|
return row.plan or "", row.role or "user"
|
|
|
|
|
|
def model_allowed_for_user(model_id: str, user_id: UUID) -> bool:
|
|
"""该用户(按 plan/role 档位)能否使用 model_id;非抛出版本,供降级判断。"""
|
|
from core.model_access import is_allowed
|
|
|
|
plan, role = user_plan_role(user_id)
|
|
return is_allowed(model_id, plan, role)
|
|
|
|
|
|
def assert_model_allowed(model_id: str, user_id: Optional[UUID], kind: str) -> None:
|
|
"""档位门控(显式选择点用):user_id 非空时校验该用户能否用 model_id,不许 → 403。
|
|
|
|
用于"用户主动选模型"的入口(建 task 带 profile / 切模型 / 发媒体)——前端下拉已按档过滤,
|
|
选到档外只可能是构造请求,直接 403(纵深防御)。
|
|
与之相对,"老 task 下次发消息"走 _downgrade 静默落回 flash(见 send / optimize),不报错。
|
|
user_id=None 的内部路径(定时任务执行)不门控。
|
|
"""
|
|
if user_id is None:
|
|
return
|
|
if not model_allowed_for_user(model_id, user_id):
|
|
raise HTTPException(403, f"模型 {model_id!r} 未对你的账户开放({kind});请联系管理员调整档位")
|
|
|
|
|
|
def resolve_model_profile(profile: str, user_id: Optional[UUID] = None) -> tuple[str, str]:
|
|
"""校验 model_profile 并返回 (profile, model_id)。
|
|
|
|
传空 → cfg["default_model"]。profile 走 ModelCapabilities.load:
|
|
格式或文件错误一律 400。返 (profile_str, caps.model_id) —— 调 ensure_local_task_row
|
|
时 model_profile / model 两列一起填,保持现有 schema 双列约定。
|
|
user_id 非空 → 额外过档位门控(assert_model_allowed),档外模型 403。
|
|
"""
|
|
from core.agent_builder import load_config
|
|
from core.capabilities import ModelCapabilities
|
|
from core.paths import ROOT
|
|
|
|
cfg = load_config()
|
|
name = (profile or "").strip() or cfg["default_model"]
|
|
try:
|
|
caps = ModelCapabilities.load(name, ROOT / cfg["models_dir"])
|
|
except (FileNotFoundError, ValueError) as e:
|
|
raise HTTPException(400, f"invalid model_profile {name!r}: {e}")
|
|
assert_model_allowed(name, user_id, "text")
|
|
return name, caps.model_id
|
|
|
|
|
|
def skill_pinned_profiles() -> set:
|
|
"""全部内置 skill 定向的 model profile 集合(send 侧档外降级豁免用)。
|
|
|
|
task 落到档外模型只有两条路:管理员下调档位(该降),或 skill 定向(load_skill
|
|
热切写入,不该降 —— 产品决策放行)。档内模型根本进不了降级分支,所以
|
|
"profile ∈ 本集合"即可判定 skill 定向来源,无需追溯该 task 当初怎么切的。
|
|
只扫内置来源(与热切 switcher 的"只信内置"一致);registry 现扫 ~3ms 且只在
|
|
"当前模型已档外"的罕见分支才调。frontmatter 删掉 `model:` 行 → 集合随之收缩,
|
|
存量定向 task 下条消息自然落回 flash。
|
|
"""
|
|
from core.agent_builder import load_config
|
|
from core.paths import ROOT
|
|
from core.skills import SkillRegistry, SkillSource
|
|
|
|
cfg = load_config()
|
|
reg = SkillRegistry(SkillSource(ROOT / cfg.get("skills_dir", "skills"), "builtin"))
|
|
return {s.model for s in reg.skills.values() if s.model}
|
|
|
|
|
|
def list_media_variants(kind: str) -> list[tuple[str, dict]]:
|
|
"""扫 config/media/*.yaml 的 <kind> 段 → [(variant_key, variant_cfg), ...]。
|
|
|
|
多 provider(doubao / unifyllm)合并列举,文件名排序保证豆包在前(默认 variant
|
|
的选择逻辑另见 default_image_variant)。variant key 约定跨文件唯一;万一撞名先读的生效(后者跳过)。
|
|
目录不存在或段空 / 仅注释 → 返 []。不要求对应 API key 已设 —— 仅纯元数据列举,
|
|
UI 拉这个画下拉。真正调用时 agent_builder 那边再过 `ArkConfig.load()`
|
|
(没 key → tool 不注册)。
|
|
"""
|
|
from core.paths import ROOT
|
|
import yaml as _yaml
|
|
|
|
media_dir = ROOT / "config" / "media"
|
|
if not media_dir.is_dir():
|
|
return []
|
|
out: list[tuple[str, dict]] = []
|
|
seen: set[str] = set()
|
|
for p in sorted(media_dir.glob("*.yaml")):
|
|
try:
|
|
data = _yaml.safe_load(p.read_text(encoding="utf-8")) or {}
|
|
except Exception:
|
|
continue
|
|
for k, v in (data.get(kind) or {}).items():
|
|
if isinstance(v, dict) and k not in seen:
|
|
seen.add(k)
|
|
out.append((k, v))
|
|
return out
|
|
|
|
|
|
def list_image_variants() -> list[tuple[str, dict]]:
|
|
"""图像 variant 清单(跨 provider,见 list_media_variants)。"""
|
|
return list_media_variants("image")
|
|
|
|
|
|
def default_image_variant(allowed: Optional[set]) -> str:
|
|
"""该用户(档位过滤后)的默认 image variant key;无可用 → ""。
|
|
|
|
有 gpt_image 权限(pro 档 / admin)→ 默认 GPT 生图;否则落过滤后第一个
|
|
(= 豆包 seedream,沿用 yaml 文件名排序)。用户仍可在顶栏下拉自行切换。
|
|
/v1/image_models 的 is_default 与 resolve_image_model 的空串解析共用这里,
|
|
保证"下拉默认选中"和"实际起 run 用的"永远一致。
|
|
"""
|
|
keys = [k for k, _ in list_image_variants() if allowed is None or k in allowed]
|
|
if not keys:
|
|
return ""
|
|
return "gpt_image" if "gpt_image" in keys else keys[0]
|
|
|
|
|
|
def resolve_image_model(variant: str, user_id: Optional[UUID] = None) -> str:
|
|
"""校验 image_model variant key。
|
|
|
|
传空 = 客户端未显式选 → 解析成该用户的默认 variant(与 /v1/image_models 的
|
|
is_default 同源,见 default_image_variant);无 user_id 的内部路径仍返空
|
|
(agent_builder fallback 到第一个 variant)。传非空 → 必须存在于
|
|
config/media/*.yaml 的 image 段,否则 400。user_id 非空 → 额外过档位门控。
|
|
"""
|
|
name = (variant or "").strip()
|
|
if not name:
|
|
if user_id is None:
|
|
return ""
|
|
from core.model_access import allowed_set
|
|
plan, role = user_plan_role(user_id)
|
|
return default_image_variant(allowed_set(plan, role))
|
|
variants = {k for k, _ in list_image_variants()}
|
|
if name not in variants:
|
|
raise HTTPException(400, f"invalid image_model {name!r}; available: {sorted(variants)}")
|
|
assert_model_allowed(name, user_id, "image")
|
|
return name
|
|
|
|
|
|
def list_video_variants() -> list[tuple[str, dict]]:
|
|
"""视频 variant 清单(跨 provider,见 list_media_variants);空 → UI 隐藏下拉。"""
|
|
return list_media_variants("video")
|
|
|
|
|
|
def resolve_video_model(variant: str, user_id: Optional[UUID] = None) -> str:
|
|
"""校验 video_model variant key(同 resolve_image_model 范式)。"""
|
|
name = (variant or "").strip()
|
|
if not name:
|
|
return ""
|
|
variants = {k for k, _ in list_video_variants()}
|
|
if name not in variants:
|
|
raise HTTPException(400, f"invalid video_model {name!r}; available: {sorted(variants)}")
|
|
assert_model_allowed(name, user_id, "video")
|
|
return name
|