Files
yingqing/core/backend/apps/ai/model_routing.py
T
2026-08-31 12:18:03 +08:00

333 lines
14 KiB
Python

"""模型能力契约与动态 Fallback 候选解析。
本模块只负责“某模型是否满足本次调用”与“失败后候选如何排序”,不执行 Provider
请求、重试、日志或账务。业务入口后续只声明 :class:`ModelRequirements`,不得按模型名
复制能力判断。
"""
from __future__ import annotations
from dataclasses import dataclass, field
from datetime import datetime
from typing import Any
from apps.ai.generation_errors import provider_error_rule_errors
from apps.ai.models import ModelConfig, ModelProvider
from apps.ai.routing_policy import load_model_routing_policy
@dataclass(frozen=True, slots=True)
class ModelRequirements:
"""一次模型调用的最小能力需求;未使用的字段保持空值。"""
capability: str
operation: str
features: frozenset[str] = field(default_factory=frozenset)
reference_mode: str | None = None
reference_images: int = 0
reference_videos: int = 0
reference_audios: int = 0
aspect_ratio: str | None = None
resolution: str | None = None
duration: int | None = None
language: str | None = None
public_voice: str | None = None
char_count: int | None = None
speed_ratio: float | None = None
output_format: str | None = None
def __post_init__(self) -> None:
if not self.capability or not self.operation:
raise ValueError("模型能力需求必须包含 capability 和 operation")
for name in ("reference_images", "reference_videos", "reference_audios"):
if getattr(self, name) < 0:
raise ValueError(f"{name} 不能为负数")
if self.reference_mode not in {None, "none", "single", "multiple"}:
raise ValueError("reference_mode 只能是 none、single 或 multiple")
if self.reference_mode == "single" and self.reference_images != 1:
raise ValueError("reference_mode=single 时 reference_images 必须等于 1")
if self.reference_mode == "multiple" and self.reference_images < 2:
raise ValueError("reference_mode=multiple 时 reference_images 必须至少为 2")
# Seedance 的视频编辑模式用 -1 表示“跟随参考视频时长”。
if self.duration is not None and self.duration <= 0 and self.duration != -1:
raise ValueError("duration 必须大于 0,或为 -1(跟随参考视频)")
if self.char_count is not None and self.char_count < 0:
raise ValueError("char_count 不能为负数")
if self.speed_ratio is not None and self.speed_ratio <= 0:
raise ValueError("speed_ratio 必须大于 0")
@dataclass(frozen=True, slots=True)
class CapabilityMatch:
matched: bool
reasons: tuple[str, ...] = ()
def _dict(value: Any) -> dict[str, Any]:
return value if isinstance(value, dict) else {}
def _set(value: Any) -> set[str]:
if not isinstance(value, (list, tuple, set, frozenset)):
return set()
return {str(item) for item in value if str(item)}
def _positive_int(value: Any) -> int | None:
if isinstance(value, bool):
return None
try:
result = int(value)
except (TypeError, ValueError):
return None
return result if result >= 0 else None
def routing_metadata(model: ModelConfig) -> dict[str, Any]:
return _dict(_dict(model.metadata).get("routing"))
def capability_metadata(model: ModelConfig) -> dict[str, Any]:
return _dict(_dict(model.metadata).get("capabilities"))
def model_allows_fallback(model: ModelConfig) -> bool:
"""该模型失败后是否允许向外切换;缺失配置时保守关闭。"""
return routing_metadata(model).get("fallback_on_failure") is True
def model_is_fallback_candidate(model: ModelConfig) -> bool:
"""该模型是否允许被其他失败任务选中;缺失配置时保守关闭。"""
return routing_metadata(model).get("fallback_candidate") is True
def provider_fallback_priority(provider: ModelProvider) -> int:
"""读取供应商候选优先级;数字越小越优先,缺失或非法值统一排到普通供应商层。"""
value = _dict(_dict(provider.metadata).get("routing")).get("fallback_priority", 100)
if isinstance(value, bool):
return 100
try:
return int(value)
except (TypeError, ValueError):
return 100
def provider_metadata_errors(metadata: Any) -> tuple[str, ...]:
"""校验供应商路由与错误展示配置;供后台写入校验和运行前诊断共同复用。"""
routing = _dict(_dict(metadata).get("routing"))
errors = list(provider_error_rule_errors(metadata))
if "fallback_priority" in routing:
value = routing["fallback_priority"]
if isinstance(value, bool) or not isinstance(value, int) or not 0 <= value <= 1000:
errors.append("routing.fallback_priority 必须是 0 到 1000 的整数,数字越小越优先")
return tuple(errors)
def model_metadata_errors(capability: str, metadata: Any) -> tuple[str, ...]:
"""校验模型路由开关与最小能力契约,缺少候选开关时允许保存。"""
root = _dict(metadata)
routing = _dict(root.get("routing"))
errors: list[str] = []
for name in ("fallback_on_failure", "fallback_candidate"):
if name in routing and not isinstance(routing[name], bool):
errors.append(f"routing.{name} 必须是 true 或 false")
if routing.get("fallback_candidate") is not True:
return tuple(errors)
capabilities = _dict(root.get("capabilities"))
if not _set(capabilities.get("operations")):
errors.append("capabilities.operations 至少配置一个操作")
if capability == ModelConfig.Capability.IMAGE:
modes = _set(capabilities.get("reference_modes"))
if modes - {"none", "single", "multiple"}:
errors.append("capabilities.reference_modes 只能包含 none、single、multiple")
if capability == ModelConfig.Capability.VIDEO:
if not _set(capabilities.get("resolutions")):
errors.append("视频候选必须配置 capabilities.resolutions")
durations = capabilities.get("durations")
if not isinstance(durations, (list, tuple)) or not durations:
errors.append("视频候选必须配置 capabilities.durations")
if capability == ModelConfig.Capability.AUDIO:
voice_map = capabilities.get("voice_map")
if voice_map is not None and not isinstance(voice_map, dict):
errors.append("配音 capabilities.voice_map 必须是公开音色到供应商音色 ID 的对象")
return tuple(errors)
def capability_metadata_errors(model: ModelConfig) -> tuple[str, ...]:
"""返回会导致模型无法安全进入候选池的中文配置问题。"""
return model_metadata_errors(model.capability, model.metadata)
def _check_limit(
reasons: list[str],
capabilities: dict[str, Any],
field_name: str,
required: int,
label: str,
) -> None:
if required <= 0:
return
maximum = _positive_int(capabilities.get(field_name))
if maximum is None or maximum < required:
reasons.append(f"{label}上限不足:需要 {required},配置为 {maximum if maximum is not None else '缺失'}")
def _video_pricing_available(model: ModelConfig, resolution: str | None) -> bool:
if not resolution:
return True
pricing = _dict(_dict(model.metadata).get("pricing"))
if not pricing:
return False
# 1080p/4k 必须有精确价格档;480p/720p 可使用现有 default 档。
tier = pricing.get(resolution) if resolution in {"1080p", "4k"} else pricing.get(resolution) or pricing.get("default")
return isinstance(tier, dict) and bool(tier)
def match_model_requirements(model: ModelConfig, requirements: ModelRequirements) -> CapabilityMatch:
"""纯函数式能力匹配;缺失配置一律保守排除,不按模型名猜测。"""
reasons: list[str] = []
if model.capability != requirements.capability:
reasons.append(f"能力大类不匹配:需要 {requirements.capability},模型为 {model.capability}")
return CapabilityMatch(False, tuple(reasons))
capabilities = capability_metadata(model)
operations = _set(capabilities.get("operations"))
if requirements.operation not in operations:
reasons.append(f"不支持操作 {requirements.operation}")
missing_features = sorted(set(requirements.features) - _set(capabilities.get("features")))
if missing_features:
reasons.append(f"缺少特性:{', '.join(missing_features)}")
if requirements.reference_mode and requirements.reference_mode != "none":
if requirements.reference_mode not in _set(capabilities.get("reference_modes")):
reasons.append(f"不支持 {requirements.reference_mode} 参考图模式")
_check_limit(reasons, capabilities, "max_reference_images", requirements.reference_images, "参考图片")
_check_limit(reasons, capabilities, "max_reference_videos", requirements.reference_videos, "参考视频")
_check_limit(reasons, capabilities, "max_reference_audios", requirements.reference_audios, "参考音频")
for field_name, required, label in (
("aspect_ratios", requirements.aspect_ratio, "画面比例"),
("resolutions", requirements.resolution, "分辨率"),
("languages", requirements.language, "语言"),
("output_formats", requirements.output_format, "输出格式"),
):
# 视频编辑的 adaptive 比例由火山按输入视频决定,不是模型静态能力表中的比例档。
if required and required != "adaptive" and required not in _set(capabilities.get(field_name)):
reasons.append(f"不支持{label} {required}")
if requirements.duration is not None and requirements.duration != -1:
durations = capabilities.get("durations")
supported = set(durations) if isinstance(durations, (list, tuple, set, frozenset)) else set()
if requirements.duration not in supported:
reasons.append(f"不支持时长 {requirements.duration} 秒")
if requirements.public_voice:
voice_map = capabilities.get("voice_map")
if not isinstance(voice_map, dict) or not voice_map.get(requirements.public_voice):
reasons.append(f"缺少公开音色 {requirements.public_voice} 的供应商映射")
if requirements.char_count is not None:
max_chars = _positive_int(capabilities.get("max_chars"))
if max_chars is None or requirements.char_count > max_chars:
reasons.append(
f"字符上限不足:需要 {requirements.char_count},配置为 {max_chars if max_chars is not None else '缺失'}"
)
if requirements.speed_ratio is not None:
speed_range = capabilities.get("speed_range")
if not isinstance(speed_range, (list, tuple)) or len(speed_range) != 2:
reasons.append("缺少合法的语速范围配置")
else:
try:
minimum, maximum = float(speed_range[0]), float(speed_range[1])
except (TypeError, ValueError):
reasons.append("语速范围配置不是数字")
else:
if not minimum <= requirements.speed_ratio <= maximum:
reasons.append(f"不支持语速 {requirements.speed_ratio:g}")
if requirements.capability == ModelConfig.Capability.VIDEO and not _video_pricing_available(
model, requirements.resolution
):
reasons.append(f"缺少分辨率 {requirements.resolution} 的视频价格配置")
return CapabilityMatch(not reasons, tuple(reasons))
def _timestamp(value: datetime | None) -> float:
return value.timestamp() if value is not None else 0.0
def _candidate_sort_key(model: ModelConfig) -> tuple[int, float, float, str]:
return (
provider_fallback_priority(model.provider),
-_timestamp(model.updated_at),
-_timestamp(model.created_at),
str(model.id),
)
def resolve_fallback_candidates(
*,
primary_model: ModelConfig,
requirements: ModelRequirements,
attempted_model_ids: set[Any] | frozenset[Any] = frozenset(),
excluded_provider_ids: set[Any] | frozenset[Any] = frozenset(),
) -> list[ModelConfig]:
"""在主模型明确失败后解析候选;不会替换或重排用户传入的主模型。"""
if not model_allows_fallback(primary_model):
return []
attempted = {str(value) for value in attempted_model_ids}
attempted.add(str(primary_model.id))
excluded_providers = {str(value) for value in excluded_provider_ids}
policy = load_model_routing_policy()
remaining_model_slots = max(0, policy.max_models - len(attempted))
if remaining_model_slots == 0:
return []
candidates = []
queryset = ModelConfig.objects.select_related("provider").filter(
capability=requirements.capability,
status=ModelConfig.Status.ACTIVE,
provider__status=ModelProvider.Status.ACTIVE,
)
for model in queryset:
if str(model.id) in attempted or str(model.provider_id) in excluded_providers:
continue
if not model_is_fallback_candidate(model) or capability_metadata_errors(model):
continue
if match_model_requirements(model, requirements).matched:
candidates.append(model)
candidates.sort(key=_candidate_sort_key)
return candidates[:remaining_model_slots]
__all__ = [
"CapabilityMatch",
"ModelRequirements",
"capability_metadata",
"capability_metadata_errors",
"match_model_requirements",
"model_allows_fallback",
"model_is_fallback_candidate",
"model_metadata_errors",
"provider_fallback_priority",
"provider_metadata_errors",
"resolve_fallback_candidates",
"routing_metadata",
]