feat: 接入模型动态 Fallback 与调用审计

This commit is contained in:
hh
2026-07-21 10:35:50 +08:00
parent 9429133d04
commit d0c3690d47
45 changed files with 9253 additions and 173 deletions
+53
View File
@@ -48,6 +48,59 @@ ASSETS_API_PROJECT_NAME=default
# 豆包语音合成(旁白配音 TTS)· 火山控制台-语音技术-语音合成
VOLC_TTS_APPID=8945759494
VOLC_TTS_ACCESS_TOKEN=w7Ye8FdTHADU05PV5cVNjud8FseOnYzR
# 模型失败重试与动态 Fallback 统一策略。修改后必须同时重启 Django API 与 Celery Worker。
# 单个逻辑任务最多尝试的不同模型数(包含用户选择或系统默认的主模型)。
MODEL_ROUTING_MAX_MODELS=3
# 单个逻辑任务最多发出的真实模型请求数(包含首次、重试和 Fallback)。
MODEL_ROUTING_MAX_CALLS=5
# 重试退避的随机抖动比例;0.20 表示基础等待时间上下浮动 20%。
MODEL_ROUTING_JITTER_RATIO=0.20
# 文本:原模型失败后的重试等待秒数;逗号分隔,列表长度就是重试次数。
MODEL_ROUTING_TEXT_RETRY_DELAYS=1,3
# 文本:普通非流式请求单次超时秒数。
MODEL_ROUTING_TEXT_REQUEST_TIMEOUT=120
# 文本:脚本等流式请求单次超时秒数。
MODEL_ROUTING_TEXT_STREAM_TIMEOUT=300
# 文本:一次逻辑任务从首次请求起的总时限秒数。
MODEL_ROUTING_TEXT_TOTAL_TIMEOUT=480
# 文本:429 Retry-After 允许等待的最长秒数。
MODEL_ROUTING_TEXT_RETRY_AFTER_CAP=15
# 图片:原模型失败后的重试等待秒数;默认重试 1 次。
MODEL_ROUTING_IMAGE_RETRY_DELAYS=3
# 图片:单次生成或编辑请求超时秒数。
MODEL_ROUTING_IMAGE_REQUEST_TIMEOUT=300
# 图片:单张图片逻辑任务总时限秒数。
MODEL_ROUTING_IMAGE_TOTAL_TIMEOUT=900
# 图片:429 Retry-After 允许等待的最长秒数。
MODEL_ROUTING_IMAGE_RETRY_AFTER_CAP=15
# 配音:原模型失败后的重试等待秒数;默认重试 1 次。
MODEL_ROUTING_AUDIO_RETRY_DELAYS=2
# 配音:单次 TTS 请求超时秒数。
MODEL_ROUTING_AUDIO_REQUEST_TIMEOUT=60
# 配音:一次配音逻辑任务总时限秒数。
MODEL_ROUTING_AUDIO_TOTAL_TIMEOUT=180
# 配音:429 Retry-After 允许等待的最长秒数。
MODEL_ROUTING_AUDIO_RETRY_AFTER_CAP=15
# 视频:尚未拿到远端任务 ID 时的提交重试等待秒数;状态未知时禁止重提。
MODEL_ROUTING_VIDEO_SUBMIT_RETRY_DELAYS=3
# 视频:单次创建远端任务的请求超时秒数。
MODEL_ROUTING_VIDEO_SUBMIT_TIMEOUT=120
# 视频:提交、重试和 Fallback 阶段的总时限秒数。
MODEL_ROUTING_VIDEO_SUBMIT_TOTAL_TIMEOUT=300
# 视频:单次查询远端任务状态的超时秒数。
MODEL_ROUTING_VIDEO_POLL_REQUEST_TIMEOUT=60
# 视频:成功提交后允许等待成片的最长秒数。
MODEL_ROUTING_VIDEO_GENERATION_TIMEOUT=1800
# 视频:429 Retry-After 允许等待的最长秒数。
MODEL_ROUTING_VIDEO_RETRY_AFTER_CAP=30
# 后处理:下载、上传和资产保存的重试等待秒数;不会重新调用模型。
MODEL_ROUTING_POSTPROCESS_RETRY_DELAYS=2,5,10
# 模特库三视图真实生成与计费开关:true 允许预留积分并投递 worker;false 时提交接口直接拒绝任务。
MODEL_TRIVIEW_GENERATION_ENABLED=true
+55
View File
@@ -37,3 +37,58 @@ MODEL_TRIVIEW_GENERATION_ENABLED=true
# false + 空白名单 = 全部使用旧提示词。修改后 API 与 Worker 必须使用同一组值并同步重启。
MODEL_TRYON_PROMPT_V2_ENABLED=false
MODEL_TRYON_PROMPT_V2_CANARY_TEAM_IDS=
# 模型失败重试与自动 Fallback 统一策略。修改后需同步重启 API 与 Celery Worker。
# 所有时间值单位均为秒;重试等待时间使用英文逗号分隔,留空表示该阶段不重试。
# max_models 包含用户选择的主模型;max_calls 包含首次调用、原模型重试和候选模型调用。
# 单个逻辑任务最多尝试的不同模型数量,范围 1~10。调大可能增加等待时间与平台成本。
MODEL_ROUTING_MAX_MODELS=3
# 单个逻辑任务最多发出的真实模型请求总数,范围 1~20,且不得小于 MAX_MODELS。
MODEL_ROUTING_MAX_CALLS=5
# 重试等待的随机抖动比例,范围 0~1;0.20 表示基础等待时间随机浮动 ±20%。
MODEL_ROUTING_JITTER_RATIO=0.20
# 文本:默认原模型最多重试 2 次,分别等待 1 秒、3 秒。
MODEL_ROUTING_TEXT_RETRY_DELAYS=1,3
# 普通文本单次调用超时;调小可能截断正常生成,调大会延长故障等待。
MODEL_ROUTING_TEXT_REQUEST_TIMEOUT=120
# 流式文本单次调用超时;脚本 SSE 通常比普通文本耗时更长。
MODEL_ROUTING_TEXT_STREAM_TIMEOUT=300
# 单个文本逻辑任务总时限,默认 480 秒(8 分钟)。
MODEL_ROUTING_TEXT_TOTAL_TIMEOUT=480
# 文本遇到 429 时接受 Retry-After 的最长等待时间。
MODEL_ROUTING_TEXT_RETRY_AFTER_CAP=15
# 图片:默认原模型等待 3 秒后重试 1 次。
MODEL_ROUTING_IMAGE_RETRY_DELAYS=3
# 单次图片生成或编辑请求超时;中转站真实出图可能超过 75 秒。
MODEL_ROUTING_IMAGE_REQUEST_TIMEOUT=300
# 单张图片逻辑任务总时限,默认 900 秒(15 分钟)。
MODEL_ROUTING_IMAGE_TOTAL_TIMEOUT=900
# 图片遇到 429 时接受 Retry-After 的最长等待时间。
MODEL_ROUTING_IMAGE_RETRY_AFTER_CAP=15
# 配音:默认原模型等待 2 秒后重试 1 次。
MODEL_ROUTING_AUDIO_RETRY_DELAYS=2
# 单次配音请求超时。
MODEL_ROUTING_AUDIO_REQUEST_TIMEOUT=60
# 单个配音逻辑任务总时限,默认 180 秒(3 分钟)。
MODEL_ROUTING_AUDIO_TOTAL_TIMEOUT=180
# 配音遇到 429 时接受 Retry-After 的最长等待时间。
MODEL_ROUTING_AUDIO_RETRY_AFTER_CAP=15
# 视频提交:仅在尚未取得 Provider 任务 ID 时,等待 3 秒后重试 1 次。
MODEL_ROUTING_VIDEO_SUBMIT_RETRY_DELAYS=3
# 单次创建视频 Provider 任务的超时。
MODEL_ROUTING_VIDEO_SUBMIT_TIMEOUT=120
# 视频提交与 Fallback 阶段总时限,默认 300 秒(5 分钟)。
MODEL_ROUTING_VIDEO_SUBMIT_TOTAL_TIMEOUT=300
# 单次查询视频 Provider 任务状态的超时。
MODEL_ROUTING_VIDEO_POLL_REQUEST_TIMEOUT=60
# 成功提交后最长成片等待时间,默认 1800 秒(30 分钟);超时不代表可以重复提交。
MODEL_ROUTING_VIDEO_GENERATION_TIMEOUT=1800
# 视频提交遇到 429 时接受 Retry-After 的最长等待时间。
MODEL_ROUTING_VIDEO_RETRY_AFTER_CAP=30
# 模型成功后的下载、上传、资产保存重试间隔;不增加模型调用次数,也不触发 Fallback。
MODEL_ROUTING_POSTPROCESS_RETRY_DELAYS=2,5,10
+7 -2
View File
@@ -1,8 +1,8 @@
# AirShelf 后端架构
> 适用目录:`core/backend`。
> 最后核对:2026-07-18
<!-- architecture-sync-commit: bfca22677cd0637b151c801fe53e68eb4ed8a6e3 -->
> 最后核对:2026-07-21。
<!-- architecture-sync-commit: 9429133d04e695e53ba5c485e3e13a379fb3a03a -->
> 本文描述当前代码结构;系统级拓扑见 [core/ARCHITECTURE.md](../ARCHITECTURE.md)。
## 1. 定位
@@ -60,6 +60,8 @@ core/backend/
- SQLite/MySQL 数据库切换。
- Redis cache、Celery broker/result 和业务锁。
- TOS、火山 ARK、火山 TTS、YunQi、TokenSSR 配置。
- `MODEL_ROUTING_POLICY` 集中声明文本、图片、配音和视频的重试、单次超时、总时限、最多候选模型数、最多真实调用数及后处理退避。
- 模型路由策略允许由环境变量覆盖;修改后需要同步重启 API 与 Celery Worker,保证两端使用相同策略。
- AI 功能开关和灰度团队配置。
- 积分赠送等平台参数。
@@ -184,6 +186,7 @@ API 总入口:
- `ModelProvider`
- `ModelConfig`
- `AITask`
- `AIModelAttempt`
- `ImageConversation`
- `QualityWord`
- `PromptTemplate`
@@ -192,6 +195,8 @@ API 总入口:
- 实体提取、基础资产、三视图、生图、生视频和 TTS
- 用户错误转换和生成通知
`AITask` 表示一次用户逻辑任务及其唯一计费生命周期;它与 `AIModelAttempt` 是一对多关系。每条 `AIModelAttempt` 对应一次真实模型请求,按顺序保存供应商/模型快照、重试或 Fallback 标记、状态、耗时、错误、用量、平台成本及脱敏请求/响应摘要,不独立参与积分预留、扣费或任务状态推进。
主要目录和文件:
```text
+102
View File
@@ -24,6 +24,41 @@ def env_list(name: str, default: str = "") -> list[str]:
return [item.strip() for item in value.split(",") if item.strip()]
def env_int(name: str, default: int) -> int:
"""读取整数环境变量;格式错误时给出可直接定位配置项的中文错误。"""
value = os.getenv(name)
if value is None or not value.strip():
return default
try:
return int(value)
except ValueError as exc:
raise ValueError(f"环境变量 {name} 必须是整数,当前值:{value!r}") from exc
def env_float(name: str, default: float) -> float:
"""读取小数环境变量;格式错误时给出可直接定位配置项的中文错误。"""
value = os.getenv(name)
if value is None or not value.strip():
return default
try:
return float(value)
except ValueError as exc:
raise ValueError(f"环境变量 {name} 必须是数字,当前值:{value!r}") from exc
def env_int_list(name: str, default: tuple[int, ...]) -> list[int]:
"""读取逗号分隔的整数列表;空字符串表示不重试。"""
value = os.getenv(name)
if value is None:
return list(default)
if not value.strip():
return []
try:
return [int(item.strip()) for item in value.split(",") if item.strip()]
except ValueError as exc:
raise ValueError(f"环境变量 {name} 必须是逗号分隔的整数列表,当前值:{value!r}") from exc
SECRET_KEY = env("DJANGO_SECRET_KEY", "airshelf-dev-insecure-key")
DEBUG = env_bool("DJANGO_DEBUG", False)
ALLOWED_HOSTS = env_list("DJANGO_ALLOWED_HOSTS", "localhost,127.0.0.1")
@@ -235,6 +270,73 @@ PROVIDER_API_VERSIONS = {
# 自由创作视频:团队在途任务并发上限(视频长时高价,限并发同时压 web 慢调用敞口)
FREE_VIDEO_MAX_CONCURRENT = int(env("FREE_VIDEO_MAX_CONCURRENT", "3"))
# 模型失败重试与动态 Fallback 的统一策略配置,已由全部 AI 业务入口与调用审计共用。
#
# 生效方式:修改环境变量后,需要同步重启 Django API 与 Celery Worker,确保两端读取同一策略。
# 时间单位:除 jitter_ratio 外,所有 timeout、cap 和 retry_delays 均为“秒”。
# 次数口径:max_calls 包含首次调用、原模型重试和所有候选模型调用;max_models 包含主模型。
MODEL_ROUTING_POLICY = {
# 单个用户逻辑任务最多尝试的不同模型数量;包含用户选择或系统默认的主模型。
# 合法范围 1~10。调大可提高成功率,但会增加等待时间和平台真实成本。
"max_models": env_int("MODEL_ROUTING_MAX_MODELS", 3),
# 单个用户逻辑任务最多发出的真实模型请求总数;包含首次调用、重试和 Fallback。
# 合法范围 120,且不得小于 max_models。用户侧仍只允许一次预留和一次最终结算。
"max_calls": env_int("MODEL_ROUTING_MAX_CALLS", 5),
# 退避时间随机抖动比例;0.20 表示在基础等待时间上随机浮动 ±20%。
# 合法范围 0~1。调大可减轻并发任务同时重试,但会让实际等待时间波动更明显。
"jitter_ratio": env_float("MODEL_ROUTING_JITTER_RATIO", 0.20),
"text": {
# 文本主模型失败后的重试等待时间;列表长度就是重试次数,默认 1 秒、3 秒(最多重试 2 次)。
"retry_delays": env_int_list("MODEL_ROUTING_TEXT_RETRY_DELAYS", (1, 3)),
# 普通非流式文本请求的单次最长等待时间。
"request_timeout": env_int("MODEL_ROUTING_TEXT_REQUEST_TIMEOUT", 120),
# 流式文本请求的单次最长等待时间;脚本 SSE 正常生成可能明显慢于普通文本。
"stream_timeout": env_int("MODEL_ROUTING_TEXT_STREAM_TIMEOUT", 300),
# 单个文本逻辑任务从首次调用 Provider 起允许的总执行时间,默认 8 分钟。
"total_timeout": env_int("MODEL_ROUTING_TEXT_TOTAL_TIMEOUT", 480),
# 429 响应 Retry-After 的最长接受时间;超过该值时不继续长时间等待,转入 Fallback 或最终失败。
"retry_after_cap": env_int("MODEL_ROUTING_TEXT_RETRY_AFTER_CAP", 15),
},
"image": {
# 图片主模型失败后的重试等待时间;默认等待 3 秒后重试 1 次。
"retry_delays": env_int_list("MODEL_ROUTING_IMAGE_RETRY_DELAYS", (3,)),
# 单次图片生成或编辑请求的最长等待时间;中转站真实出图可能超过 75 秒。
"request_timeout": env_int("MODEL_ROUTING_IMAGE_REQUEST_TIMEOUT", 300),
# 单张图片逻辑任务的总执行时间,默认 15 分钟。
"total_timeout": env_int("MODEL_ROUTING_IMAGE_TOTAL_TIMEOUT", 900),
# 图片 429 Retry-After 的最长接受时间。
"retry_after_cap": env_int("MODEL_ROUTING_IMAGE_RETRY_AFTER_CAP", 15),
},
"audio": {
# 配音主模型失败后的重试等待时间;默认等待 2 秒后重试 1 次。
"retry_delays": env_int_list("MODEL_ROUTING_AUDIO_RETRY_DELAYS", (2,)),
# 单次配音请求的最长等待时间。
"request_timeout": env_int("MODEL_ROUTING_AUDIO_REQUEST_TIMEOUT", 60),
# 单个配音逻辑任务的总执行时间,默认 3 分钟。
"total_timeout": env_int("MODEL_ROUTING_AUDIO_TOTAL_TIMEOUT", 180),
# 配音 429 Retry-After 的最长接受时间。
"retry_after_cap": env_int("MODEL_ROUTING_AUDIO_RETRY_AFTER_CAP", 15),
},
"video": {
# 视频尚未拿到 Provider 任务 ID 时的提交重试等待时间;默认等待 3 秒后重试 1 次。
"submit_retry_delays": env_int_list("MODEL_ROUTING_VIDEO_SUBMIT_RETRY_DELAYS", (3,)),
# 单次创建视频 Provider 任务的最长等待时间。
"submit_timeout": env_int("MODEL_ROUTING_VIDEO_SUBMIT_TIMEOUT", 120),
# 视频“提交与 Fallback”阶段总时限,默认 5 分钟;成功拿到 Provider 任务 ID 后结束该阶段。
"submit_total_timeout": env_int("MODEL_ROUTING_VIDEO_SUBMIT_TOTAL_TIMEOUT", 300),
# 单次查询视频任务状态的最长等待时间。
"poll_request_timeout": env_int("MODEL_ROUTING_VIDEO_POLL_REQUEST_TIMEOUT", 60),
# 视频成功提交后的最长成片等待时间,默认 30 分钟;状态未知或普通轮询超时不得重复提交视频。
"generation_timeout": env_int("MODEL_ROUTING_VIDEO_GENERATION_TIMEOUT", 1800),
# 视频提交遇到 429 时,Retry-After 的最长接受时间。
"retry_after_cap": env_int("MODEL_ROUTING_VIDEO_RETRY_AFTER_CAP", 30),
},
"postprocess": {
# 模型成功后的下载、上传和资产保存重试等待时间;不增加模型调用次数,也不触发模型 Fallback。
"retry_delays": env_int_list("MODEL_ROUTING_POSTPROCESS_RETRY_DELAYS", (2, 5, 10)),
},
}
# 模特库三视图真实生成与计费开关:必须在 API 与 Celery worker 同步部署后才显式开启。
# 关闭时提交接口直接拒绝;开启后允许预留积分并把生成任务投递给 worker。
MODEL_TRIVIEW_GENERATION_ENABLED = env_bool("MODEL_TRIVIEW_GENERATION_ENABLED", False)
+34 -2
View File
@@ -3,7 +3,8 @@ from decimal import Decimal
from rest_framework import serializers
from apps.accounts.models import Team, TeamMember, User
from apps.ai.models import AITask, ModelConfig, ModelProvider, PromptTemplate, QualityWord
from apps.ai.model_routing import model_metadata_errors, provider_metadata_errors
from apps.ai.models import AIModelAttempt, AITask, ModelConfig, ModelProvider, PromptTemplate, QualityWord
from apps.assets.models import Asset
from apps.billing.models import CreditLedger, QuotaPolicy
from apps.projects.models import Project
@@ -162,11 +163,26 @@ class AdminTaskSerializer(serializers.ModelSerializer):
return str((actual / rate - base).quantize(Decimal("0.01")))
class AdminModelAttemptSerializer(serializers.ModelSerializer):
class Meta:
model = AIModelAttempt
fields = [
"id", "sequence", "provider_name", "provider_display_name", "model_name", "model_display_name",
"public_model_name", "capability", "operation", "status", "is_retry", "is_fallback",
"previous_attempt", "provider_task_id", "started_at", "finished_at", "duration_ms", "error_type",
"provider_error_code", "raw_error", "safe_error_summary", "usage", "platform_cost", "request_summary",
"response_summary",
]
read_only_fields = fields
class AdminTaskDetailSerializer(AdminTaskSerializer):
attempts = AdminModelAttemptSerializer(source="model_attempts", many=True, read_only=True)
class Meta(AdminTaskSerializer.Meta):
fields = AdminTaskSerializer.Meta.fields + [
"project", "idempotency_key", "request_payload", "response_payload",
"error_message", "submitted_at", "completed_at",
"error_message", "submitted_at", "completed_at", "attempts",
]
read_only_fields = fields
@@ -211,6 +227,12 @@ class AdminModelProviderSerializer(serializers.ModelSerializer):
anno = getattr(obj, "model_count_anno", None)
return anno if anno is not None else obj.models.count()
def validate_metadata(self, value):
errors = provider_metadata_errors(value)
if errors:
raise serializers.ValidationError(list(errors))
return value
class AdminModelConfigSerializer(serializers.ModelSerializer):
provider_name = serializers.CharField(source="provider.name", read_only=True, default=None)
@@ -224,6 +246,16 @@ class AdminModelConfigSerializer(serializers.ModelSerializer):
# is_default 只经 set-default 端点改,不在普通编辑里直接写
read_only_fields = ["id", "provider_name", "is_default", "created_at"]
def validate(self, attrs):
attrs = super().validate(attrs)
if self.instance is None or "metadata" in attrs or "capability" in attrs:
capability = attrs.get("capability", getattr(self.instance, "capability", ""))
metadata = attrs.get("metadata", getattr(self.instance, "metadata", {}))
errors = model_metadata_errors(capability, metadata)
if errors:
raise serializers.ValidationError({"metadata": list(errors)})
return attrs
class AdminProjectSerializer(serializers.ModelSerializer):
team_name = serializers.CharField(source="team.name", read_only=True, default=None)
+53
View File
@@ -368,6 +368,35 @@ class AdminTaskMonitorTests(TestCase):
self.assertEqual(r.status_code, 200)
self.assertGreaterEqual(r.data["count"], 4)
def test_list_defers_large_payload_columns_but_detail_keeps_them(self):
from django.db import connection
from django.test.utils import CaptureQueriesContext
self.t_ok.request_payload = {"prompt": "x" * 100_000}
self.t_ok.response_payload = {"image": "y" * 100_000}
self.t_ok.error_message = "z" * 100_000
self.t_ok.save(update_fields=["request_payload", "response_payload", "error_message", "updated_at"])
with CaptureQueriesContext(connection) as captured:
response = self.ac.get("/api/admin/tasks/?page_size=10")
self.assertEqual(response.status_code, 200)
task_selects = [
item["sql"]
for item in captured.captured_queries
if "FROM \"ai_aitask\"" in item["sql"] and "COUNT(" not in item["sql"]
]
self.assertTrue(task_selects)
for sql in task_selects:
self.assertNotIn("request_payload", sql)
self.assertNotIn("response_payload", sql)
self.assertNotIn("error_message", sql)
detail = self.ac.get(f"/api/admin/tasks/{self.t_ok.id}/")
self.assertEqual(detail.status_code, 200)
self.assertEqual(detail.data["request_payload"], self.t_ok.request_payload)
self.assertEqual(detail.data["response_payload"], self.t_ok.response_payload)
self.assertEqual(detail.data["error_message"], self.t_ok.error_message)
def test_filter_status_type_anomaly(self):
self.assertTrue(all(t["status"] == "failed" for t in self.ac.get("/api/admin/tasks/?status=failed").data["results"]))
self.assertTrue(all(t["task_type"] == "person_image" for t in self.ac.get("/api/admin/tasks/?task_type=person_image").data["results"]))
@@ -380,10 +409,34 @@ class AdminTaskMonitorTests(TestCase):
self.assertFalse(self.ac.get(f"/api/admin/tasks/{self.t_ok.id}/").data["cost_anomaly"])
def test_detail_has_payloads(self):
from django.utils import timezone
from apps.ai.models import AIModelAttempt
AIModelAttempt.objects.create(
task=self.t_failed,
sequence=1,
provider=self.mc.provider,
model_config=self.mc,
provider_name=self.mc.provider.name,
provider_display_name=self.mc.provider.display_name,
model_name=self.mc.name,
model_display_name=self.mc.display_name,
public_model_name="AirShelf Image",
capability="image",
operation="image_generate",
status=AIModelAttempt.Status.FAILED,
started_at=timezone.now(),
error_type="provider_unavailable",
)
d = self.ac.get(f"/api/admin/tasks/{self.t_failed.id}/")
self.assertEqual(d.status_code, 200)
self.assertIn("request_payload", d.data)
self.assertIn("response_payload", d.data)
self.assertEqual(len(d.data["attempts"]), 1)
self.assertEqual(d.data["attempts"][0]["model_name"], self.mc.name)
row = next(item for item in self.ac.get("/api/admin/tasks/").data["results"] if item["id"] == str(self.t_failed.id))
self.assertNotIn("attempts", row)
def test_retry_dispatch_and_audit(self):
with patch("apps.ai.tasks.generate_standalone_image_task.delay") as mock_delay:
+14 -2
View File
@@ -417,7 +417,13 @@ def admin_asset_reviews_poll(request):
@permission_classes([IsPlatformAdmin])
def admin_tasks(request):
"""全局 AITask 列表(?status= / ?task_type= / ?team= / ?anomaly=1 成本异常 筛 + 分页)。"""
qs = AITask.objects.select_related("team", "model_config").order_by("-created_at")
# 列表不返回请求/响应/完整错误正文。部分图片、视频任务的 JSON 可达数 MB,若随列表页
# 从远程 MySQL 读取,会让只有 10 行的分页请求也长时间卡在“加载中”;详情接口仍完整读取。
qs = (
AITask.objects.select_related("team", "model_config")
.defer("request_payload", "response_payload", "error_message")
.order_by("-created_at")
)
st = request.query_params.get("status")
if st in dict(AITask.Status.choices):
qs = qs.filter(status=st)
@@ -437,7 +443,13 @@ def admin_tasks(request):
@api_view(["GET"])
@permission_classes([IsPlatformAdmin])
def admin_task_detail(request, task_id):
task = AITask.objects.select_related("team", "model_config").filter(id=task_id).first()
# 尝试链只在详情请求加载,列表仍保持原查询与一任务一行。
task = (
AITask.objects.select_related("team", "model_config")
.prefetch_related("model_attempts")
.filter(id=task_id)
.first()
)
if task is None:
return Response({"detail": "not found"}, status=status.HTTP_404_NOT_FOUND)
return Response(AdminTaskDetailSerializer(task).data)
+11
View File
@@ -1,7 +1,18 @@
from django.apps import AppConfig
from django.core.exceptions import ImproperlyConfigured
class AiConfig(AppConfig):
default_auto_field = "django.db.models.BigAutoField"
name = "apps.ai"
def ready(self) -> None:
# API 与 Celery Worker 都会在 Django 启动时经过这里。路由策略一旦配置错误就立即失败,
# 避免两端带着不同或不可控的超时、重试规则继续运行。
from .routing_policy import RoutingPolicyConfigurationError, load_model_routing_policy
try:
load_model_routing_policy()
except RoutingPolicyConfigurationError as exc:
raise ImproperlyConfigured(f"模型路由策略配置错误:{exc}") from exc
+84 -1
View File
@@ -2,12 +2,15 @@ VOLCANO_PROVIDER = {
"name": "volcengine",
"display_name": "火山引擎(豆包)",
"base_url": "https://ark.cn-beijing.volces.com/api/v3",
# Fallback 候选排序:数字越小越优先。运行时只读配置,不识别供应商名称。
"metadata": {"routing": {"fallback_priority": 10}},
}
YUNQI_PROVIDER = {
"name": "yunqi",
"display_name": "YunQi AI(New API 网关)",
"base_url": "https://www.yunqiai.chat/v1",
"metadata": {"routing": {"fallback_priority": 20}},
}
YUNQI_MODELS = [
@@ -16,7 +19,12 @@ YUNQI_MODELS = [
"name": "gpt-image-2",
"capability": "image",
"endpoint": "images/generations",
"metadata": {"modes": ["text"], "response_format": "b64_json", "source": "https://www.yunqiai.chat"},
"metadata": {
"modes": ["text", "singleImage", "multiReference"],
"supports_reference": True,
"response_format": "b64_json",
"source": "https://www.yunqiai.chat",
},
},
]
@@ -132,3 +140,78 @@ VOLCANO_MODELS = [
},
]
IMAGE_RATIOS = ["1:1", "3:4", "4:5", "9:16", "16:9"]
VIDEO_RATIOS = ["21:9", "16:9", "4:3", "1:1", "3:4", "9:16"]
def _mode_limit(modes: set[str], prefix: str) -> int:
for mode in modes:
if mode.startswith(f"{prefix}:"):
try:
return int(mode.split(":", 1)[1])
except (TypeError, ValueError):
return 0
return 0
def _catalog_capabilities(item: dict) -> dict:
"""由模型目录声明生成最小能力契约;只用于写库初始化,不参与运行时猜测。"""
capability = item["capability"]
metadata = item["metadata"]
if capability == "text":
return {"operations": ["chat"], "features": ["streaming", "structured_output"]}
if capability == "image":
modes = set(metadata.get("modes") or [])
supports_reference = bool(metadata.get("supports_reference")) or bool(
modes & {"singleImage", "multiReference"}
)
return {
"operations": ["image_generate", "image_edit"] if supports_reference else ["image_generate"],
"features": [],
"reference_modes": ["none", "single", "multiple"] if supports_reference else ["none"],
"max_reference_images": 9 if supports_reference else 0,
"aspect_ratios": IMAGE_RATIOS,
}
if capability == "video":
modes = set(metadata.get("modes") or [])
features = ["text_to_video"]
for mode, feature in (
("startFrameOptional", "start_frame"),
("lastFrameOptional", "last_frame"),
("imageReference:9", "image_reference"),
("videoReference:3", "video_reference"),
("audioReference:3", "audio_reference"),
):
if mode in modes:
features.append(feature)
if metadata.get("audio") in {"optional", "required", True}:
features.append("generate_audio")
return {
"operations": ["video_generate"],
"features": features,
"max_reference_images": _mode_limit(modes, "imageReference"),
"max_reference_videos": _mode_limit(modes, "videoReference"),
"max_reference_audios": _mode_limit(modes, "audioReference"),
"aspect_ratios": VIDEO_RATIOS,
"resolutions": list(metadata.get("resolutions") or []),
"durations": list(metadata.get("durations") or []),
}
return {}
def _enrich_catalog_routing(items: list[dict], *, fallback_on_failure: bool) -> None:
for item in items:
metadata = item["metadata"]
metadata["routing"] = {
"fallback_on_failure": fallback_on_failure,
"fallback_candidate": True,
}
metadata["capabilities"] = _catalog_capabilities(item)
# 火山当前只作为候选、不向外切;YunQi 可在失败后切换。以后调整只改数据库 metadata。
_enrich_catalog_routing(VOLCANO_MODELS, fallback_on_failure=False)
_enrich_catalog_routing(YUNQI_MODELS, fallback_on_failure=True)
+89 -18
View File
@@ -15,6 +15,7 @@ import logging
import re
import uuid
from datetime import timedelta
from decimal import Decimal
from io import BytesIO
from django.conf import settings
@@ -272,18 +273,33 @@ def _reap_stale_free_video_tasks(*, team) -> None:
· SUBMITTED/POLLING 超 2 小时:轮询链早已断且无人认领(正常出片 5-10 分钟)→ 标失败退费;
· POSTPROCESSING 超 30 分钟:转存/结算中途崩溃 → 标失败退费(火山可能已出片,平台承担该笔成本)。"""
now = timezone.now()
from .routing_policy import load_model_routing_policy
video_policy = load_model_routing_policy().video
buckets = [
([AITask.Status.RESERVED], now - timedelta(minutes=10), "任务未在预期时间内提交(自动回收)"),
([AITask.Status.SUBMITTED, AITask.Status.POLLING], now - timedelta(hours=2), "生成超时(自动回收)"),
([AITask.Status.POSTPROCESSING], now - timedelta(minutes=30), "视频结果处理超时(自动回收)"),
(
[AITask.Status.RESERVED],
{"updated_at__lt": now - timedelta(seconds=video_policy.submit_total_timeout)},
"任务未在配置的提交总时限内完成(自动回收)",
),
(
[AITask.Status.SUBMITTED, AITask.Status.POLLING],
{"submitted_at__lt": now - timedelta(seconds=video_policy.generation_timeout)},
"生成超过配置的成片等待总时限(自动回收)",
),
(
[AITask.Status.POSTPROCESSING],
{"updated_at__lt": now - timedelta(minutes=30)},
"视频结果处理超时(自动回收)",
),
]
for statuses, cutoff, reason in buckets:
for statuses, stale_filter, reason in buckets:
stale = AITask.objects.filter(
team=team,
project__isnull=True,
task_type=AITask.Type.FREE_VIDEO,
status__in=statuses,
updated_at__lt=cutoff,
**stale_filter,
)
for task in stale:
try:
@@ -404,6 +420,7 @@ def submit_free_video(*, team, user, params: dict) -> AITask:
"price_multiplier": quote.meta.get("price_multiplier", "1"),
"points_per_yuan_snapshot": quote.meta.get("rate", ""),
"references": built["snapshots"],
"model_routing_v1": True,
}
# 建任务 + 预留同一事务:余额不足/限额拦截时回滚任务行,不留半套
@@ -418,7 +435,7 @@ def submit_free_video(*, team, user, params: dict) -> AITask:
idempotency_key=f"free_video:{team.id}:{uuid.uuid4()}",
request_payload=request_payload,
estimated_cost=quote.points,
base_cost=quote.base_cost_yuan,
base_cost=Decimal("0"),
)
try:
reserve_credit(team=team, user=user, task=task, amount=reserve_amount)
@@ -431,26 +448,44 @@ def submit_free_video(*, team, user, params: dict) -> AITask:
# 火山调用在事务外(不持锁调外网)
try:
from .services import build_provider
from .services import execute_routed_video_submit
provider = build_provider(model_config)
response = provider.create_video_task(
model=model_config.name,
endpoint=model_config.endpoint,
routed = execute_routed_video_submit(
task=task,
primary_model=model_config,
prompt=built["api_prompt"],
ratio=aspect_ratio,
duration=duration,
resolution=resolution,
generate_audio=generate_audio,
reference_images=[],
content_items=built["content_items"],
pricing_references=built["snapshots"],
generate_audio=generate_audio,
seed=seed if seed != -1 else None,
search_mode=search_mode,
request_summary={"feature": "free_video", "mode": mode},
)
task.provider_task_id = str(response.get("id") or response.get("task_id") or "")
response, provider_task_id = routed.value
# execute_model_call 以原子 F 表达式累计实际尝试的平台成本;刷新内存对象,
# 保证本方法返回值与数据库中的 AITask.base_cost 完全一致。
task.refresh_from_db(fields=["base_cost"])
task.provider_task_id = provider_task_id
task.response_payload = response
payload = dict(task.request_payload or {})
payload["actual_model_config_id"] = str(routed.actual_model.id)
task.request_payload = payload
task.status = AITask.Status.SUBMITTED
task.submitted_at = timezone.now()
task.save(update_fields=["provider_task_id", "response_payload", "status", "submitted_at", "updated_at"])
task.save(
update_fields=[
"provider_task_id",
"response_payload",
"request_payload",
"status",
"submitted_at",
"updated_at",
]
)
except Exception as exc: # noqa: BLE001 — 创建失败:标失败退费,返回失败卡(不向上抛)
code, raw_message = parse_provider_error(exc)
public_error = classify_generation_error(
@@ -577,11 +612,47 @@ def finalize_free_video(*, task: AITask) -> AITask:
if not task.provider_task_id:
return task
from .services import build_provider
from .routing_policy import load_model_routing_policy
from .services import get_video_provider
provider = build_provider(task.model_config)
video_policy = load_model_routing_policy().video
if task.submitted_at and (
timezone.now() - task.submitted_at
).total_seconds() >= video_policy.generation_timeout:
timeout_message = "视频生成超过配置的成片等待总时限"
public_error = classify_generation_error(
TimeoutError(timeout_message),
operation="video_generate",
reference_id=str(task.id),
)
with transaction.atomic():
locked = AITask.objects.select_for_update().get(id=task.id)
if locked.status not in (AITask.Status.SUBMITTED, AITask.Status.POLLING):
return locked
locked.status = AITask.Status.FAILED
locked.error_code = "GenerationTimeout"
locked.error_message = timeout_message
locked.completed_at = timezone.now()
locked.save(
update_fields=["status", "error_code", "error_message", "completed_at", "updated_at"]
)
release_credit(reservation=locked.credit_reservation, reason=timeout_message)
_notify_failure(locked, raw=timeout_message, hint=public_error.fallback_message)
return locked
submit_attempt = (
task.model_attempts.filter(status="succeeded", operation="video_generate")
.select_related("model_config__provider")
.order_by("-sequence")
.first()
)
actual_model = submit_attempt.model_config if submit_attempt and submit_attempt.model_config else task.model_config
provider = get_video_provider(actual_model)
response = provider.poll_video_task(
endpoint=task.model_config.endpoint, provider_task_id=task.provider_task_id
endpoint=actual_model.endpoint,
provider_task_id=task.provider_task_id,
timeout=video_policy.poll_request_timeout,
)
remote_status = str(response.get("status") or "")
@@ -647,7 +718,7 @@ def finalize_free_video(*, task: AITask) -> AITask:
from decimal import Decimal
settle = quote_video_actual(
locked.model_config, tokens=total_tokens, with_video_ref=with_video_ref, resolution=resolution,
actual_model, tokens=total_tokens, with_video_ref=with_video_ref, resolution=resolution,
multiplier=Decimal(str(payload.get("price_multiplier") or "1")),
)
actual, base_cost = settle.points, settle.base_cost_yuan
@@ -0,0 +1,154 @@
"""为存量供应商和模型补齐动态 Fallback 路由与最小能力配置。
运行时解析器不识别供应商名或模型名这里仅把当前产品已经确认的火山/YunQi优先级
写入数据库 metadata以后新增供应商或模型只维护 metadata不需要修改路由代码
"""
from django.db import migrations
DIRECT_PROVIDER_NAMES = {"volcengine", "volcano", "ark", "volcano_ark", "doubao"}
IMAGE_RATIOS = ["1:1", "3:4", "4:5", "9:16", "16:9"]
VIDEO_RATIOS = ["21:9", "16:9", "4:3", "1:1", "3:4", "9:16"]
VOICE_MAP = {
"BV700_streaming": "BV700_streaming",
"BV034_streaming": "BV034_streaming",
"BV001_streaming": "BV001_streaming",
"BV056_streaming": "BV056_streaming",
"BV102_streaming": "BV102_streaming",
"BV002_streaming": "BV002_streaming",
}
def _provider_priority(name):
if name in DIRECT_PROVIDER_NAMES:
return 10
if name == "yunqi" or name.startswith("yunqi_"):
return 20
return 100
def _mode_limit(modes, prefix, default=0):
for mode in modes:
if mode.startswith(prefix):
_, _, raw = mode.partition(":")
try:
return int(raw)
except (TypeError, ValueError):
return default
return default
def _capabilities(model):
metadata = dict(model.metadata or {})
capability = model.capability
if capability == "text":
return {
"operations": ["chat"],
"features": ["streaming", "structured_output"],
}
if capability == "image":
modes = {str(item) for item in metadata.get("modes") or []}
supports_reference = bool(metadata.get("supports_reference")) or bool(
modes & {"singleImage", "multiReference"}
)
reference_modes = ["none"]
if supports_reference or "singleImage" in modes:
reference_modes.append("single")
if supports_reference or "multiReference" in modes:
reference_modes.append("multiple")
operations = ["image_generate"]
if supports_reference:
operations.append("image_edit")
return {
"operations": operations,
"features": [],
"reference_modes": reference_modes,
"max_reference_images": 9 if supports_reference else 0,
"aspect_ratios": IMAGE_RATIOS,
}
if capability == "video":
modes = {str(item) for item in metadata.get("modes") or []}
features = ["text_to_video"]
for mode, feature in (
("startFrameOptional", "start_frame"),
("lastFrameOptional", "last_frame"),
("imageReference:9", "image_reference"),
("videoReference:3", "video_reference"),
("audioReference:3", "audio_reference"),
):
if mode in modes:
features.append(feature)
if metadata.get("audio") in {"optional", "required", True}:
features.append("generate_audio")
return {
"operations": ["video_generate"],
"features": features,
"max_reference_images": _mode_limit(modes, "imageReference", 0),
"max_reference_videos": _mode_limit(modes, "videoReference", 0),
"max_reference_audios": _mode_limit(modes, "audioReference", 0),
"aspect_ratios": VIDEO_RATIOS,
"resolutions": list(metadata.get("resolutions") or []),
"durations": list(metadata.get("durations") or []),
}
if capability == "audio":
return {
"operations": ["tts"],
"features": [],
"languages": ["zh-CN"],
"voice_map": VOICE_MAP,
"max_chars": 10000,
"speed_range": [0.5, 2.0],
"output_formats": ["mp3"],
}
return {}
def apply(apps, schema_editor):
ModelProvider = apps.get_model("ai", "ModelProvider")
ModelConfig = apps.get_model("ai", "ModelConfig")
for provider in ModelProvider.objects.all().iterator():
metadata = dict(provider.metadata or {})
routing = dict(metadata.get("routing") or {})
routing.setdefault("fallback_priority", _provider_priority(provider.name))
metadata["routing"] = routing
provider.metadata = metadata
provider.save(update_fields=["metadata"])
for model in ModelConfig.objects.select_related("provider").all().iterator():
capabilities = _capabilities(model)
metadata = dict(model.metadata or {})
routing = dict(metadata.get("routing") or {})
is_supported = bool(capabilities)
routing.setdefault("fallback_candidate", is_supported)
routing.setdefault(
"fallback_on_failure",
is_supported and model.provider.name not in DIRECT_PROVIDER_NAMES,
)
metadata["routing"] = routing
if capabilities:
metadata.setdefault("capabilities", capabilities)
model.metadata = metadata
model.save(update_fields=["metadata"])
def revert(apps, schema_editor):
ModelProvider = apps.get_model("ai", "ModelProvider")
ModelConfig = apps.get_model("ai", "ModelConfig")
for provider in ModelProvider.objects.all().iterator():
metadata = dict(provider.metadata or {})
metadata.pop("routing", None)
provider.metadata = metadata
provider.save(update_fields=["metadata"])
for model in ModelConfig.objects.all().iterator():
metadata = dict(model.metadata or {})
metadata.pop("routing", None)
metadata.pop("capabilities", None)
model.metadata = metadata
model.save(update_fields=["metadata"])
class Migration(migrations.Migration):
dependencies = [("ai", "0027_aitask_model_triview_type")]
operations = [migrations.RunPython(apply, revert)]
@@ -0,0 +1,91 @@
from django.db import migrations, models
import django.db.models.deletion
import uuid
class Migration(migrations.Migration):
dependencies = [("ai", "0028_seed_model_routing_metadata")]
operations = [
migrations.CreateModel(
name="AIModelAttempt",
fields=[
("id", models.UUIDField(default=uuid.uuid4, editable=False, primary_key=True, serialize=False)),
("created_at", models.DateTimeField(auto_now_add=True)),
("updated_at", models.DateTimeField(auto_now=True)),
("sequence", models.PositiveIntegerField()),
("provider_name", models.CharField(max_length=128)),
("provider_display_name", models.CharField(blank=True, max_length=128)),
("model_name", models.CharField(max_length=128)),
("model_display_name", models.CharField(blank=True, max_length=128)),
("public_model_name", models.CharField(blank=True, max_length=128)),
("capability", models.CharField(max_length=32)),
("operation", models.CharField(max_length=64)),
(
"status",
models.CharField(
choices=[("started", "Started"), ("succeeded", "Succeeded"), ("failed", "Failed")],
default="started",
max_length=24,
),
),
("is_retry", models.BooleanField(default=False)),
("is_fallback", models.BooleanField(default=False)),
("provider_task_id", models.CharField(blank=True, max_length=255)),
("started_at", models.DateTimeField()),
("finished_at", models.DateTimeField(blank=True, null=True)),
("duration_ms", models.PositiveBigIntegerField(blank=True, null=True)),
("error_type", models.CharField(blank=True, max_length=64)),
("provider_error_code", models.CharField(blank=True, max_length=128)),
("raw_error", models.TextField(blank=True)),
("safe_error_summary", models.TextField(blank=True)),
("usage", models.JSONField(blank=True, default=dict)),
("platform_cost", models.DecimalField(decimal_places=4, default=0, max_digits=12)),
("request_summary", models.JSONField(blank=True, default=dict)),
("response_summary", models.JSONField(blank=True, default=dict)),
(
"model_config",
models.ForeignKey(
blank=True,
null=True,
on_delete=django.db.models.deletion.SET_NULL,
to="ai.modelconfig",
),
),
(
"previous_attempt",
models.ForeignKey(
blank=True,
null=True,
on_delete=django.db.models.deletion.SET_NULL,
related_name="next_attempts",
to="ai.aimodelattempt",
),
),
(
"provider",
models.ForeignKey(
blank=True,
null=True,
on_delete=django.db.models.deletion.SET_NULL,
to="ai.modelprovider",
),
),
(
"task",
models.ForeignKey(
on_delete=django.db.models.deletion.CASCADE,
related_name="model_attempts",
to="ai.aitask",
),
),
],
options={
"ordering": ["sequence"],
"indexes": [models.Index(fields=["task", "status"], name="ai_attempt_task_status_idx")],
"constraints": [
models.UniqueConstraint(fields=("task", "sequence"), name="ai_attempt_task_sequence_unique")
],
},
),
]
+329
View File
@@ -0,0 +1,329 @@
"""模型能力契约与动态 Fallback 候选解析。
本模块只负责某模型是否满足本次调用失败后候选如何排序不执行 Provider
请求重试日志或账务业务入口后续只声明 :class:`ModelRequirements`不得按模型名
复制能力判断
"""
from __future__ import annotations
from dataclasses import dataclass, field
from datetime import datetime
from typing import Any
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")
if self.duration is not None and self.duration <= 0:
raise ValueError("duration 必须大于 0")
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"))
if "fallback_priority" not in routing:
return ()
value = routing["fallback_priority"]
if isinstance(value, bool) or not isinstance(value, int) or not 0 <= value <= 1000:
return ("routing.fallback_priority 必须是 0 到 1000 的整数,数字越小越优先",)
return ()
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, "输出格式"),
):
if required and required not in _set(capabilities.get(field_name)):
reasons.append(f"不支持{label} {required}")
if requirements.duration is not None:
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",
]
+54
View File
@@ -174,6 +174,60 @@ class AITask(TeamOwnedModel):
return f"{self.task_type}:{self.status}:{self.id}"
class AIModelAttempt(TimeStampedModel):
"""AITask 下的一次真实模型请求审计;不参与积分预留、扣费或任务生命周期。"""
class Status(models.TextChoices):
STARTED = "started", "Started"
SUCCEEDED = "succeeded", "Succeeded"
FAILED = "failed", "Failed"
task = models.ForeignKey(AITask, on_delete=models.CASCADE, related_name="model_attempts")
sequence = models.PositiveIntegerField()
provider = models.ForeignKey(ModelProvider, on_delete=models.SET_NULL, null=True, blank=True)
model_config = models.ForeignKey(ModelConfig, on_delete=models.SET_NULL, null=True, blank=True)
# 名称均保存历史快照,后续后台改名或删除配置也不影响旧任务审计。
provider_name = models.CharField(max_length=128)
provider_display_name = models.CharField(max_length=128, blank=True)
model_name = models.CharField(max_length=128)
model_display_name = models.CharField(max_length=128, blank=True)
public_model_name = models.CharField(max_length=128, blank=True)
capability = models.CharField(max_length=32)
operation = models.CharField(max_length=64)
status = models.CharField(max_length=24, choices=Status.choices, default=Status.STARTED)
is_retry = models.BooleanField(default=False)
is_fallback = models.BooleanField(default=False)
previous_attempt = models.ForeignKey(
"self",
on_delete=models.SET_NULL,
null=True,
blank=True,
related_name="next_attempts",
)
provider_task_id = models.CharField(max_length=255, blank=True)
started_at = models.DateTimeField()
finished_at = models.DateTimeField(null=True, blank=True)
duration_ms = models.PositiveBigIntegerField(null=True, blank=True)
error_type = models.CharField(max_length=64, blank=True)
provider_error_code = models.CharField(max_length=128, blank=True)
raw_error = models.TextField(blank=True)
safe_error_summary = models.TextField(blank=True)
usage = models.JSONField(default=dict, blank=True)
platform_cost = models.DecimalField(max_digits=12, decimal_places=4, default=0)
request_summary = models.JSONField(default=dict, blank=True)
response_summary = models.JSONField(default=dict, blank=True)
class Meta:
ordering = ["sequence"]
constraints = [
models.UniqueConstraint(fields=["task", "sequence"], name="ai_attempt_task_sequence_unique"),
]
indexes = [models.Index(fields=["task", "status"], name="ai_attempt_task_status_idx")]
def __str__(self) -> str:
return f"{self.task_id}:{self.sequence}:{self.status}"
class QualityWord(TimeStampedModel):
"""平台单层质量词配置:按生成阶段(stage)+ 槽位(slot)挂若干质量词,
生成侧拼提示词时优先读取本配置;**无配置则回落各 builder 的写死值**(零回归)
@@ -86,6 +86,7 @@ class OpenAICompatibleProvider(VolcanoArkProvider):
endpoint: str = "images/generations",
image: str | list[str] | None = None,
size: str = "1024x1536",
timeout: float = 300,
) -> dict[str, Any]:
"""文生图(可选单图参考 base64)。多图参考请用 image_edit。返回体含 url 或 b64_json。"""
if not self.api_key:
@@ -98,7 +99,7 @@ class OpenAICompatibleProvider(VolcanoArkProvider):
self._endpoint_url(endpoint),
headers={"Authorization": f"Bearer {self.api_key}", "Content-Type": "application/json"},
json=body,
timeout=300,
timeout=timeout,
)
_raise_with_body(response)
return response.json()
@@ -111,6 +112,7 @@ class OpenAICompatibleProvider(VolcanoArkProvider):
images: list[str],
endpoint: str = "images/edits",
size: str = "1024x1536",
timeout: float = 300,
) -> dict[str, Any]:
"""参考图编辑/合成(gpt-image-2 核心能力):multipart `image[]` 上传一张或多张参考图。
@@ -139,7 +141,47 @@ class OpenAICompatibleProvider(VolcanoArkProvider):
headers={"Authorization": f"Bearer {self.api_key}"}, # multipart 不要手设 Content-Type
files=files,
data=data,
timeout=300,
timeout=timeout,
)
_raise_with_body(response)
return response.json()
def synthesize(
self,
*,
model: str,
text: str,
voice_type: str,
speed_ratio: float = 1.0,
endpoint: str = "audio/speech",
output_format: str = "mp3",
timeout: float = 60,
uid: str = "airshelf",
) -> tuple[bytes, int]:
"""OpenAI 兼容 ``audio/speech``:返回音频字节和可选时长毫秒。
``uid`` 仅为与火山直连调用签名一致不发送给标准 OpenAI 请求新增供应商只需在
ModelConfig endpoint / voice_map无需新增 Provider 分支
"""
del uid
if not self.api_key:
raise ValueError("中转站 api_key 未配置")
response = requests.post(
self._endpoint_url(endpoint or "audio/speech"),
headers={"Authorization": f"Bearer {self.api_key}", "Content-Type": "application/json"},
json={
"model": model,
"input": text,
"voice": voice_type,
"speed": float(speed_ratio or 1.0),
"response_format": output_format,
},
timeout=timeout,
)
_raise_with_body(response)
duration_ms = 0
try:
duration_ms = int(float(response.headers.get("x-audio-duration-ms") or 0))
except (TypeError, ValueError):
duration_ms = 0
return response.content, duration_ms
+33 -9
View File
@@ -59,7 +59,14 @@ class VolcanoArkProvider:
payload=data,
)
def chat_completion(self, *, model: str, messages: list[dict[str, str]], endpoint: str = "chat/completions") -> dict[str, Any]:
def chat_completion(
self,
*,
model: str,
messages: list[dict[str, str]],
endpoint: str = "chat/completions",
timeout: float = 120,
) -> dict[str, Any]:
if not self.api_key:
raise ValueError("VOLCANO_ARK_API_KEY is not configured")
@@ -67,7 +74,7 @@ class VolcanoArkProvider:
f"{self.base_url.rstrip('/')}/{endpoint.lstrip('/')}",
headers={"Authorization": f"Bearer {self.api_key}", "Content-Type": "application/json"},
json={"model": model, "messages": messages},
timeout=120,
timeout=timeout,
)
response.raise_for_status()
return response.json()
@@ -80,6 +87,7 @@ class VolcanoArkProvider:
endpoint: str = "chat/completions",
temperature: float = 0.8,
extra_body: dict[str, Any] | None = None,
timeout: float = 300,
) -> Iterator[dict[str, Any]]:
"""流式对话:逐块 yield {type:'delta'|'tool_call'|'done', ...}。
OpenAI 兼容 SSE(火山 ARK / 各中转站同构),供脚本 agent SSE 端实时转发"""
@@ -97,7 +105,7 @@ class VolcanoArkProvider:
},
json=body,
stream=True,
timeout=300,
timeout=timeout,
) as response:
response.raise_for_status()
# SSE 响应常不带 charset,requests 会按 latin-1 解码 → 中文乱码。强制 UTF-8。
@@ -157,6 +165,7 @@ class VolcanoArkProvider:
endpoint: str = "images/generations",
image: str | list[str] | None = None,
size: str = "2K",
timeout: float = 180,
) -> dict[str, Any]:
if not self.api_key:
raise ValueError("VOLCANO_ARK_API_KEY is not configured")
@@ -174,7 +183,7 @@ class VolcanoArkProvider:
f"{self.base_url.rstrip('/')}/{endpoint.lstrip('/')}",
headers={"Authorization": f"Bearer {self.api_key}", "Content-Type": "application/json"},
json=body,
timeout=180,
timeout=timeout,
)
# 非 2xx 抽出火山真实报错(error.code/message)而非裸 HTTP 400 —— 否则 Seedream 的
# 内容审核 / 尺寸过小(历史 4:5 套图 raise_for_status 吞 body 的踩坑)等真因看不见。
@@ -195,6 +204,7 @@ class VolcanoArkProvider:
content_items: list[dict[str, Any]] | None = None,
seed: int | None = None,
search_mode: str = "off",
timeout: float = 120,
) -> dict[str, Any]:
if not self.api_key:
raise ValueError("VOLCANO_ARK_API_KEY is not configured")
@@ -223,18 +233,24 @@ class VolcanoArkProvider:
f"{self.base_url.rstrip('/')}/{endpoint.lstrip('/')}",
headers={"Authorization": f"Bearer {self.api_key}", "Content-Type": "application/json"},
json=body,
timeout=120,
timeout=timeout,
)
_raise_volcano_error(response)
return response.json()
def poll_video_task(self, *, endpoint: str, provider_task_id: str) -> dict[str, Any]:
def poll_video_task(
self,
*,
endpoint: str,
provider_task_id: str,
timeout: float = 60,
) -> dict[str, Any]:
if not self.api_key:
raise ValueError("VOLCANO_ARK_API_KEY is not configured")
response = requests.get(
f"{self.base_url.rstrip('/')}/{endpoint.rstrip('/')}/{provider_task_id}",
headers={"Authorization": f"Bearer {self.api_key}"},
timeout=60,
timeout=timeout,
)
_raise_volcano_error(response)
return response.json()
@@ -305,7 +321,15 @@ class VolcanoTtsProvider:
def configured(self) -> bool:
return bool(self.appid and self.access_token)
def synthesize(self, *, text: str, voice_type: str, speed_ratio: float = 1.0, uid: str = "airshelf") -> tuple[bytes, int]:
def synthesize(
self,
*,
text: str,
voice_type: str,
speed_ratio: float = 1.0,
uid: str = "airshelf",
timeout: float = 60,
) -> tuple[bytes, int]:
"""合成一段语音。返回 (mp3 字节, 时长毫秒;接口没回时长则为 0)。"""
if not self.configured:
raise TtsNotConfigured(
@@ -322,7 +346,7 @@ class VolcanoTtsProvider:
self.base_url,
headers={"Authorization": f"Bearer;{self.access_token}"},
json=body,
timeout=60,
timeout=timeout,
)
response.raise_for_status()
data = response.json()
+2 -1
View File
@@ -26,6 +26,7 @@ class YunqiProvider(VolcanoArkProvider):
endpoint: str = "images/generations",
image: str | list[str] | None = None,
size: str = "1024x1536",
timeout: float = 300,
) -> dict[str, Any]:
if not self.api_key:
raise ValueError("YUNQI_API_KEY is not configured")
@@ -37,7 +38,7 @@ class YunqiProvider(VolcanoArkProvider):
f"{self.base_url.rstrip('/')}/{endpoint.lstrip('/')}",
headers={"Authorization": f"Bearer {self.api_key}", "Content-Type": "application/json"},
json=body,
timeout=300,
timeout=timeout,
)
response.raise_for_status()
return response.json()
+456
View File
@@ -0,0 +1,456 @@
"""统一模型请求执行器:重试、动态切换和每次真实调用的 Feedback 留痕。"""
from __future__ import annotations
from dataclasses import dataclass, field
from decimal import Decimal, InvalidOperation
import random
import re
import time
from typing import Any, Callable, Generic, Mapping, TypeVar
from django.db import transaction
from django.db.models import F, Max
from django.utils import timezone
from apps.ai.generation_errors import PublicGenerationError, classify_generation_error
from apps.ai.model_routing import ModelRequirements, resolve_fallback_candidates
from apps.ai.models import AIModelAttempt, AITask, ModelConfig
from apps.ai.routing_policy import load_model_routing_policy
T = TypeVar("T")
@dataclass(frozen=True, slots=True)
class AttemptMetadata:
"""调用方可返回的审计摘要;不得放入密钥、鉴权头或完整素材。"""
provider_task_id: str = ""
usage: Mapping[str, Any] = field(default_factory=dict)
platform_cost: Decimal = Decimal("0")
response_summary: Mapping[str, Any] = field(default_factory=dict)
@dataclass(frozen=True, slots=True)
class RoutingExecutionResult(Generic[T]):
value: T
actual_model: ModelConfig
attempt: AIModelAttempt
call_count: int
fallback_used: bool
@dataclass(frozen=True, slots=True)
class ErrorRoutingDecision:
retry_current: bool
fallback: bool
exclude_provider: bool = False
_TERMINAL_ERRORS = {
"content_rejected",
"invalid_input",
"asset_unavailable",
"user_credit_insufficient",
"processing_failed",
}
_FALLBACK_ONLY_ERRORS = {"provider_quota_exhausted", "provider_config_error", "model_unavailable"}
_PROVIDER_SCOPE_ERRORS = {"provider_quota_exhausted", "provider_config_error"}
_SENSITIVE_KEYS = {"api_key", "apikey", "authorization", "password", "secret", "access_token", "refresh_token"}
_SECRET_PATTERNS = (
re.compile(r"(?i)(authorization\s*[:=]\s*)([^\s,;]+)"),
re.compile(r"(?i)((?:api[_-]?key|access[_-]?token|secret)\s*[:=]\s*)([^\s,;]+)"),
re.compile(r"\bsk-[A-Za-z0-9_-]{8,}\b"),
)
def decide_error_routing(error: PublicGenerationError) -> ErrorRoutingDecision:
"""把统一安全错误映射成路由动作;业务入口不得复制错误字符串判断。"""
if error.code in _TERMINAL_ERRORS:
return ErrorRoutingDecision(retry_current=False, fallback=False)
if error.code in _FALLBACK_ONLY_ERRORS:
return ErrorRoutingDecision(
retry_current=False,
fallback=True,
exclude_provider=error.code in _PROVIDER_SCOPE_ERRORS,
)
# 网络、超时、限流、5xx 及暂未精确分类的 Provider 异常:主模型按能力策略重试后再切换。
return ErrorRoutingDecision(retry_current=True, fallback=True)
def _redact_text(value: Any, limit: int = 4000) -> str:
text = str(value or "")
for pattern in _SECRET_PATTERNS:
if pattern.groups >= 2:
text = pattern.sub(r"\1[REDACTED]", text)
else:
text = pattern.sub("[REDACTED]", text)
return text[:limit]
def sanitize_attempt_summary(value: Any, *, depth: int = 0) -> Any:
"""生成可持久化的轻量脱敏摘要,避免日志保存密钥、二进制或完整长素材。"""
if depth >= 4:
return "[摘要层级已截断]"
if isinstance(value, Mapping):
result = {}
for index, (key, item) in enumerate(value.items()):
if index >= 40:
result["_truncated"] = True
break
name = str(key)
if name.lower() in _SENSITIVE_KEYS:
result[name] = "[REDACTED]"
else:
result[name] = sanitize_attempt_summary(item, depth=depth + 1)
return result
if isinstance(value, (list, tuple, set, frozenset)):
items = list(value)
return [sanitize_attempt_summary(item, depth=depth + 1) for item in items[:20]]
if isinstance(value, (bytes, bytearray, memoryview)):
return f"[二进制 {len(value)} 字节]"
if value is None or isinstance(value, (bool, int, float)):
return value
if isinstance(value, Decimal):
return str(value)
return _redact_text(value, limit=500)
def _coerce_metadata(value: Any) -> AttemptMetadata:
if isinstance(value, AttemptMetadata):
return AttemptMetadata(
provider_task_id=str(value.provider_task_id or "")[:255],
usage=value.usage if isinstance(value.usage, Mapping) else {},
platform_cost=value.platform_cost,
response_summary=value.response_summary if isinstance(value.response_summary, Mapping) else {},
)
if not isinstance(value, Mapping):
return AttemptMetadata()
raw_cost = value.get("platform_cost", 0)
try:
cost = Decimal(str(raw_cost or 0))
except (InvalidOperation, TypeError, ValueError):
cost = Decimal("0")
return AttemptMetadata(
provider_task_id=str(value.get("provider_task_id") or value.get("task_id") or "")[:255],
usage=value.get("usage") if isinstance(value.get("usage"), Mapping) else {},
platform_cost=max(cost, Decimal("0")),
response_summary=value.get("response_summary")
if isinstance(value.get("response_summary"), Mapping)
else {},
)
def _default_result_metadata(result: Any, model: ModelConfig) -> AttemptMetadata:
del model
if not isinstance(result, Mapping):
return AttemptMetadata()
provider_task_id = result.get("provider_task_id") or result.get("task_id") or result.get("id") or ""
usage = result.get("usage") if isinstance(result.get("usage"), Mapping) else {}
return AttemptMetadata(
provider_task_id=str(provider_task_id)[:255],
usage=usage,
response_summary={"provider_task_id": str(provider_task_id)[:255]} if provider_task_id else {},
)
def _default_error_metadata(exc: Exception, model: ModelConfig) -> AttemptMetadata:
del model
return _coerce_metadata(getattr(exc, "attempt_metadata", None))
def _provider_error_code(exc: Exception) -> str:
response = getattr(exc, "response", None)
if response is not None:
try:
payload = response.json()
error = payload.get("error") if isinstance(payload, Mapping) else None
if isinstance(error, Mapping):
return str(error.get("code") or error.get("type") or "")[:128]
except Exception: # noqa: BLE001 - Provider 错误响应不保证 JSON 格式。
pass
return str(getattr(exc, "code", "") or "")[:128]
@transaction.atomic
def _start_attempt(
*,
task: AITask,
model: ModelConfig,
requirements: ModelRequirements,
public_model_name: str,
is_retry: bool,
is_fallback: bool,
previous_attempt: AIModelAttempt | None,
request_summary: Mapping[str, Any],
) -> AIModelAttempt:
# 锁父任务只用于分配稳定序号,不修改父任务,因而不会刷新 AITask.updated_at。
AITask.objects.select_for_update().only("id").get(pk=task.pk)
sequence = (
AIModelAttempt.objects.filter(task_id=task.pk).aggregate(max_sequence=Max("sequence"))["max_sequence"] or 0
) + 1
provider = model.provider
return AIModelAttempt.objects.create(
task_id=task.pk,
sequence=sequence,
provider=provider,
model_config=model,
provider_name=provider.name,
provider_display_name=provider.display_name,
model_name=model.name,
model_display_name=model.display_name,
public_model_name=public_model_name,
capability=requirements.capability,
operation=requirements.operation,
is_retry=is_retry,
is_fallback=is_fallback,
previous_attempt=previous_attempt,
started_at=timezone.now(),
request_summary=sanitize_attempt_summary(request_summary),
)
def _finish_attempt(
attempt: AIModelAttempt,
*,
status: str,
elapsed_seconds: float,
metadata: AttemptMetadata,
error: PublicGenerationError | None = None,
exc: Exception | None = None,
) -> None:
cost = max(Decimal(metadata.platform_cost), Decimal("0")).quantize(Decimal("0.0001"))
attempt.status = status
attempt.finished_at = timezone.now()
attempt.duration_ms = max(0, round(elapsed_seconds * 1000))
attempt.provider_task_id = metadata.provider_task_id[:255]
attempt.usage = sanitize_attempt_summary(metadata.usage)
attempt.platform_cost = cost
attempt.response_summary = sanitize_attempt_summary(metadata.response_summary)
if error is not None:
attempt.error_type = error.code
attempt.provider_error_code = _provider_error_code(exc) if exc is not None else ""
attempt.raw_error = _redact_text(exc)
attempt.safe_error_summary = error.fallback_message
attempt.save(
update_fields=[
"status",
"finished_at",
"duration_ms",
"provider_task_id",
"usage",
"platform_cost",
"response_summary",
"error_type",
"provider_error_code",
"raw_error",
"safe_error_summary",
"updated_at",
]
)
if cost > 0:
# QuerySet.update 不触发父任务 auto_now,避免尝试日志影响僵尸任务回收时间口径。
AITask.objects.filter(pk=attempt.task_id).update(base_cost=F("base_cost") + cost)
def _retry_after_seconds(exc: Exception) -> float | None:
response = getattr(exc, "response", None)
headers = getattr(response, "headers", None)
if not isinstance(headers, Mapping):
return None
raw = headers.get("Retry-After") or headers.get("retry-after")
try:
value = float(raw)
except (TypeError, ValueError):
return None
return max(0.0, value)
def _ability_policy(requirements: ModelRequirements):
policy = load_model_routing_policy()
if requirements.capability == ModelConfig.Capability.TEXT:
request_timeout = policy.text.stream_timeout if "streaming" in requirements.features else policy.text.request_timeout
return policy, policy.text.retry_delays, request_timeout, policy.text.total_timeout, policy.text.retry_after_cap
if requirements.capability == ModelConfig.Capability.IMAGE:
return policy, policy.image.retry_delays, policy.image.request_timeout, policy.image.total_timeout, policy.image.retry_after_cap
if requirements.capability == ModelConfig.Capability.AUDIO:
return policy, policy.audio.retry_delays, policy.audio.request_timeout, policy.audio.total_timeout, policy.audio.retry_after_cap
if requirements.capability == ModelConfig.Capability.VIDEO:
return (
policy,
policy.video.submit_retry_delays,
policy.video.submit_timeout,
policy.video.submit_total_timeout,
policy.video.retry_after_cap,
)
raise ValueError(f"暂不支持能力 {requirements.capability} 的统一模型路由")
def execute_model_call(
*,
task: AITask,
primary_model: ModelConfig,
requirements: ModelRequirements,
public_model_name: str,
invoke: Callable[[ModelConfig, float], T],
request_summary: Mapping[str, Any] | Callable[[ModelConfig], Mapping[str, Any]] | None = None,
result_metadata: Callable[[T, ModelConfig], AttemptMetadata | Mapping[str, Any]] | None = None,
error_metadata: Callable[[Exception, ModelConfig], AttemptMetadata | Mapping[str, Any]] | None = None,
error_classifier: Callable[[Exception, ModelConfig], PublicGenerationError] | None = None,
candidate_resolver: Callable[..., list[ModelConfig]] | None = None,
sleep: Callable[[float], None] = time.sleep,
monotonic: Callable[[], float] = time.monotonic,
uniform: Callable[[float, float], float] = random.uniform,
) -> RoutingExecutionResult[T]:
"""在一个既有 AITask 内完成真实调用;不创建预留、不扣费、不释放积分。"""
if task.model_config_id != primary_model.id:
raise ValueError("AITask.model_config 必须保持为用户选择或系统默认的主模型")
if primary_model.capability != requirements.capability:
raise ValueError("主模型 capability 与本次 ModelRequirements 不一致")
policy, retry_delays, request_timeout, total_timeout, retry_after_cap = _ability_policy(requirements)
resolver = candidate_resolver or resolve_fallback_candidates
result_meta = result_metadata or _default_result_metadata
error_meta = error_metadata or _default_error_metadata
classify = error_classifier or (
lambda exc, model: classify_generation_error(
exc,
operation=requirements.operation,
provider_name=model.provider.name,
reference_id=str(task.id),
)
)
route_started = monotonic()
current_model = primary_model
attempted_model_ids: set[Any] = set()
excluded_provider_ids: set[Any] = set()
previous_attempt: AIModelAttempt | None = None
calls = 0
primary_retry_index = 0
fallback_used = False
last_exception: Exception | None = None
while calls < policy.max_calls:
remaining = total_timeout - (monotonic() - route_started)
if remaining <= 0:
if last_exception is not None:
raise last_exception
raise TimeoutError("模型路由总时限已耗尽")
is_retry = current_model.id == primary_model.id and primary_retry_index > 0
summary = request_summary(current_model) if callable(request_summary) else (request_summary or {})
attempt = _start_attempt(
task=task,
model=current_model,
requirements=requirements,
public_model_name=public_model_name,
is_retry=is_retry,
is_fallback=fallback_used,
previous_attempt=previous_attempt,
request_summary=summary,
)
calls += 1
call_started = monotonic()
try:
value = invoke(current_model, min(float(request_timeout), max(0.001, remaining)))
except Exception as exc:
elapsed = monotonic() - call_started
public_error = classify(exc, current_model)
decision = decide_error_routing(public_error)
try:
failed_metadata = _coerce_metadata(error_meta(exc, current_model))
except Exception as metadata_exc: # noqa: BLE001 - 审计摘要失败不能遮蔽原始 Provider 异常。
failed_metadata = AttemptMetadata(
response_summary={"metadata_error": _redact_text(metadata_exc, limit=300)}
)
# 异步视频若已经拿到 Provider 任务 ID,表示提交结果并非“明确未创建”;禁止盲目重提。
if requirements.capability == ModelConfig.Capability.VIDEO and failed_metadata.provider_task_id:
decision = ErrorRoutingDecision(retry_current=False, fallback=False)
_finish_attempt(
attempt,
status=AIModelAttempt.Status.FAILED,
elapsed_seconds=elapsed,
metadata=failed_metadata,
error=public_error,
exc=exc,
)
previous_attempt = attempt
last_exception = exc
if decision.exclude_provider:
excluded_provider_ids.add(current_model.provider_id)
if not decision.fallback:
raise
can_retry_primary = (
not fallback_used
and decision.retry_current
and primary_retry_index < len(retry_delays)
and calls < policy.max_calls
)
if can_retry_primary:
retry_after = _retry_after_seconds(exc)
if retry_after is not None and retry_after > retry_after_cap:
can_retry_primary = False
else:
base_delay = retry_after if retry_after is not None else float(retry_delays[primary_retry_index])
delay = base_delay if retry_after is not None else base_delay * (1 + uniform(-policy.jitter_ratio, policy.jitter_ratio))
if monotonic() - route_started + max(0.0, delay) >= total_timeout:
can_retry_primary = False
else:
sleep(max(0.0, delay))
primary_retry_index += 1
if can_retry_primary:
continue
attempted_model_ids.add(current_model.id)
if calls >= policy.max_calls:
raise
candidates = resolver(
primary_model=current_model,
requirements=requirements,
attempted_model_ids=attempted_model_ids,
excluded_provider_ids=excluded_provider_ids,
)
if not candidates:
raise
current_model = candidates[0]
fallback_used = True
primary_retry_index = 0
continue
try:
metadata = _coerce_metadata(result_meta(value, current_model))
except Exception as metadata_exc: # noqa: BLE001 - 内容已生成,不能因审计摘要失败重新调用模型。
metadata = AttemptMetadata(response_summary={"metadata_error": _redact_text(metadata_exc, limit=300)})
_finish_attempt(
attempt,
status=AIModelAttempt.Status.SUCCEEDED,
elapsed_seconds=monotonic() - call_started,
metadata=metadata,
)
return RoutingExecutionResult(
value=value,
actual_model=current_model,
attempt=attempt,
call_count=calls,
fallback_used=fallback_used,
)
if last_exception is not None:
raise last_exception
raise TimeoutError("模型路由调用次数已耗尽")
__all__ = [
"AttemptMetadata",
"ErrorRoutingDecision",
"RoutingExecutionResult",
"decide_error_routing",
"execute_model_call",
"sanitize_attempt_summary",
]
+303
View File
@@ -0,0 +1,303 @@
"""模型路由策略的集中读取与校验。
本模块只负责把 Django settings 中的 ``MODEL_ROUTING_POLICY`` 转成不可变配置对象
不执行重试FallbackProvider 调用或账务操作业务入口后续只消费这里返回的策略
不得各自维护 timeout / sleep / attempts 常量
"""
from __future__ import annotations
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from typing import Any
from django.conf import settings
class RoutingPolicyConfigurationError(ValueError):
"""模型路由策略格式或取值不合法。"""
@dataclass(frozen=True, slots=True)
class TextRoutingPolicy:
retry_delays: tuple[float, ...]
request_timeout: float
stream_timeout: float
total_timeout: float
retry_after_cap: float
@dataclass(frozen=True, slots=True)
class RequestRoutingPolicy:
retry_delays: tuple[float, ...]
request_timeout: float
total_timeout: float
retry_after_cap: float
@dataclass(frozen=True, slots=True)
class VideoRoutingPolicy:
submit_retry_delays: tuple[float, ...]
submit_timeout: float
submit_total_timeout: float
poll_request_timeout: float
generation_timeout: float
retry_after_cap: float
@dataclass(frozen=True, slots=True)
class PostprocessRoutingPolicy:
retry_delays: tuple[float, ...]
@dataclass(frozen=True, slots=True)
class ModelRoutingPolicy:
max_models: int
max_calls: int
jitter_ratio: float
text: TextRoutingPolicy
image: RequestRoutingPolicy
audio: RequestRoutingPolicy
video: VideoRoutingPolicy
postprocess: PostprocessRoutingPolicy
_LABELS = {
"max_models": "最多尝试模型数量",
"max_calls": "最多真实模型调用次数",
"jitter_ratio": "重试随机抖动比例",
"retry_delays": "重试等待时间列表",
"request_timeout": "单次调用超时",
"stream_timeout": "单次流式调用超时",
"total_timeout": "逻辑任务总时限",
"retry_after_cap": "Retry-After 最长等待时间",
"submit_retry_delays": "视频提交重试等待时间列表",
"submit_timeout": "视频单次提交超时",
"submit_total_timeout": "视频提交阶段总时限",
"poll_request_timeout": "视频单次轮询超时",
"generation_timeout": "视频成片等待总时限",
}
def _error(path: str, message: str) -> RoutingPolicyConfigurationError:
key = path.rsplit(".", 1)[-1].split("[", 1)[0]
label = _LABELS.get(key, path)
return RoutingPolicyConfigurationError(f"{path}{label}{message}")
def _mapping(value: Any, path: str) -> Mapping[str, Any]:
if not isinstance(value, Mapping):
raise _error(path, f"必须是对象,当前值:{value!r}")
return value
def _required(source: Mapping[str, Any], key: str, path: str) -> Any:
if key not in source:
raise _error(f"{path}.{key}", "为必填配置")
return source[key]
def _integer(value: Any, path: str, *, minimum: int, maximum: int) -> int:
if isinstance(value, bool) or not isinstance(value, int):
raise _error(path, f"必须是整数,当前值:{value!r}")
if not minimum <= value <= maximum:
raise _error(path, f"必须在 {minimum}{maximum} 之间,当前值:{value!r}")
return value
def _number(value: Any, path: str, *, minimum: float, maximum: float) -> float:
if isinstance(value, bool) or not isinstance(value, (int, float)):
raise _error(path, f"必须是数字,当前值:{value!r}")
result = float(value)
if not minimum <= result <= maximum:
raise _error(path, f"必须在 {minimum:g}{maximum:g} 之间,当前值:{value!r}")
return result
def _delays(value: Any, path: str) -> tuple[float, ...]:
if isinstance(value, (str, bytes)) or not isinstance(value, Sequence):
raise _error(path, f"必须是数字列表,当前值:{value!r}")
return tuple(
_number(item, f"{path}[{index}]", minimum=0, maximum=3600)
for index, item in enumerate(value)
)
def _ensure_not_greater(*, smaller: float, larger: float, smaller_path: str, larger_path: str) -> None:
if smaller > larger:
raise _error(
smaller_path,
f"不得大于 {larger_path}(任务总时限),当前为 {smaller:g} > {larger:g}",
)
def _field_number(
source: Mapping[str, Any],
key: str,
path: str,
*,
minimum: float = 0,
maximum: float = 86400,
) -> float:
return _number(_required(source, key, path), f"{path}.{key}", minimum=minimum, maximum=maximum)
def _field_delays(source: Mapping[str, Any], key: str, path: str) -> tuple[float, ...]:
return _delays(_required(source, key, path), f"{path}.{key}")
def _validate_budget(
*,
total: float,
total_path: str,
bounded_values: Mapping[str, float],
delays: tuple[float, ...] = (),
delays_path: str = "",
) -> None:
for value_path, value in bounded_values.items():
_ensure_not_greater(
smaller=value,
larger=total,
smaller_path=value_path,
larger_path=total_path,
)
for index, delay in enumerate(delays):
_ensure_not_greater(
smaller=delay,
larger=total,
smaller_path=f"{delays_path}[{index}]",
larger_path=total_path,
)
def _request_policy(raw: Any, path: str) -> RequestRoutingPolicy:
source = _mapping(raw, path)
retry_delays = _field_delays(source, "retry_delays", path)
request_timeout = _field_number(source, "request_timeout", path, minimum=1)
total_timeout = _field_number(source, "total_timeout", path, minimum=1)
retry_after_cap = _field_number(source, "retry_after_cap", path, maximum=3600)
_validate_budget(
total=total_timeout,
total_path=f"{path}.total_timeout",
bounded_values={
f"{path}.request_timeout": request_timeout,
f"{path}.retry_after_cap": retry_after_cap,
},
delays=retry_delays,
delays_path=f"{path}.retry_delays",
)
return RequestRoutingPolicy(
retry_delays=retry_delays,
request_timeout=request_timeout,
total_timeout=total_timeout,
retry_after_cap=retry_after_cap,
)
def _text_policy(raw: Any) -> TextRoutingPolicy:
path = "MODEL_ROUTING_POLICY.text"
source = _mapping(raw, path)
retry_delays = _field_delays(source, "retry_delays", path)
request_timeout = _field_number(source, "request_timeout", path, minimum=1)
stream_timeout = _field_number(source, "stream_timeout", path, minimum=1)
total_timeout = _field_number(source, "total_timeout", path, minimum=1)
retry_after_cap = _field_number(source, "retry_after_cap", path, maximum=3600)
_validate_budget(
total=total_timeout,
total_path=f"{path}.total_timeout",
bounded_values={
f"{path}.request_timeout": request_timeout,
f"{path}.stream_timeout": stream_timeout,
f"{path}.retry_after_cap": retry_after_cap,
},
delays=retry_delays,
delays_path=f"{path}.retry_delays",
)
return TextRoutingPolicy(
retry_delays=retry_delays,
request_timeout=request_timeout,
stream_timeout=stream_timeout,
total_timeout=total_timeout,
retry_after_cap=retry_after_cap,
)
def _video_policy(raw: Any) -> VideoRoutingPolicy:
path = "MODEL_ROUTING_POLICY.video"
source = _mapping(raw, path)
retry_delays = _field_delays(source, "submit_retry_delays", path)
submit_timeout = _field_number(source, "submit_timeout", path, minimum=1)
submit_total_timeout = _field_number(source, "submit_total_timeout", path, minimum=1)
poll_request_timeout = _field_number(source, "poll_request_timeout", path, minimum=1)
generation_timeout = _field_number(source, "generation_timeout", path, minimum=1)
retry_after_cap = _field_number(source, "retry_after_cap", path, maximum=3600)
_validate_budget(
total=submit_total_timeout,
total_path=f"{path}.submit_total_timeout",
bounded_values={
f"{path}.submit_timeout": submit_timeout,
f"{path}.retry_after_cap": retry_after_cap,
},
delays=retry_delays,
delays_path=f"{path}.submit_retry_delays",
)
_validate_budget(
total=generation_timeout,
total_path=f"{path}.generation_timeout",
bounded_values={f"{path}.poll_request_timeout": poll_request_timeout},
)
return VideoRoutingPolicy(
submit_retry_delays=retry_delays,
submit_timeout=submit_timeout,
submit_total_timeout=submit_total_timeout,
poll_request_timeout=poll_request_timeout,
generation_timeout=generation_timeout,
retry_after_cap=retry_after_cap,
)
def load_model_routing_policy(raw_policy: Any | None = None) -> ModelRoutingPolicy:
"""读取并校验模型路由策略,返回不可变对象。
``raw_policy`` 主要供测试和启动检查注入省略时读取 Django settings
配置错误统一抛出带中文字段含义错误值和合法范围的异常
"""
raw = settings.MODEL_ROUTING_POLICY if raw_policy is None else raw_policy
source = _mapping(raw, "MODEL_ROUTING_POLICY")
path = "MODEL_ROUTING_POLICY"
max_models = _integer(_required(source, "max_models", path), f"{path}.max_models", minimum=1, maximum=10)
max_calls = _integer(_required(source, "max_calls", path), f"{path}.max_calls", minimum=1, maximum=20)
if max_calls < max_models:
raise _error(
"MODEL_ROUTING_POLICY.max_calls",
f"不得小于 max_models(最多尝试模型数量),当前为 {max_calls} < {max_models}",
)
jitter_ratio = _field_number(source, "jitter_ratio", path, maximum=1)
postprocess_path = f"{path}.postprocess"
postprocess_source = _mapping(_required(source, "postprocess", path), postprocess_path)
return ModelRoutingPolicy(
max_models=max_models,
max_calls=max_calls,
jitter_ratio=jitter_ratio,
text=_text_policy(_required(source, "text", path)),
image=_request_policy(_required(source, "image", path), f"{path}.image"),
audio=_request_policy(_required(source, "audio", path), f"{path}.audio"),
video=_video_policy(_required(source, "video", path)),
postprocess=PostprocessRoutingPolicy(
retry_delays=_field_delays(postprocess_source, "retry_delays", postprocess_path)
),
)
__all__ = [
"ModelRoutingPolicy",
"PostprocessRoutingPolicy",
"RequestRoutingPolicy",
"RoutingPolicyConfigurationError",
"TextRoutingPolicy",
"VideoRoutingPolicy",
"load_model_routing_policy",
]
+74 -20
View File
@@ -20,6 +20,7 @@ from __future__ import annotations
import json
import re
from decimal import Decimal
from functools import lru_cache
from pathlib import Path
@@ -574,7 +575,7 @@ def stream_script_agent(
):
"""生成 SSE 帧字符串的同步生成器,供 StreamingHttpResponse 包裹。
target_index 非空 = 精准只改第 N (读全脚本上下文,后端强制保留其余镜原样)"""
from apps.ai.services import build_provider, create_ai_task
from apps.ai.services import create_ai_task, stream_routed_text_request
yield _sse({"type": "tool", "id": "skill", "label": "加载电商脚本技能", "status": "running"})
skill_loaded = bool(load_ecommerce_skill())
@@ -619,6 +620,9 @@ def stream_script_agent(
"mode": mode,
"aspect_ratio": aspect_ratio,
"total_duration": total_duration,
"base_version_id": str(base_version_id or ""),
"target_index": target_index,
"model_routing_v1": True,
},
)
except Exception as exc: # noqa: BLE001 — 多为额度不足
@@ -639,13 +643,45 @@ def stream_script_agent(
task.status = AITask.Status.SUBMITTED
task.submitted_at = timezone.now()
task.save(update_fields=["status", "submitted_at", "updated_at"])
provider = build_provider(model_config)
for ev in provider.chat_completion_stream(
model=model_config.name,
endpoint=model_config.endpoint,
def validate_script_text(raw_text: str) -> dict:
candidate = normalize_draft(
raw_text,
aspect_ratio=aspect_ratio,
total_duration=effective_duration,
)
if target_index is not None and base_draft:
return _merge_single_segment(
base_draft,
candidate,
target_index,
aspect_ratio,
effective_duration,
)
return candidate
routed_stream = stream_routed_text_request(
task=task,
primary_model=model_config,
messages=messages,
streaming=True,
structured_output=True,
business_operation="script_generate",
temperature=0.85,
):
validate_text=validate_script_text,
request_summary={
"mode": mode,
"target_index": target_index,
"base_version_id": str(base_version_id or ""),
"aspect_ratio": aspect_ratio,
"total_duration": effective_duration,
},
)
while True:
try:
ev = next(routed_stream)
except StopIteration as completed:
routed = completed.value
break
et = ev.get("type")
if et == "reasoning":
# 思考流:推理模型在出 JSON 前会先想很久,把思考逐字下发给前端(像对话一样可见),
@@ -668,12 +704,8 @@ def stream_script_agent(
if piece.strip():
yield _sse({"type": "delta", "text": piece})
elif et == "done":
break
raw = "".join(full)
draft = normalize_draft(raw, aspect_ratio=aspect_ratio, total_duration=effective_duration)
if target_index is not None and base_draft:
# 精准改一镜:只采用新稿的第 target_index 镜,其余镜强制保持基准稿原样
draft = _merge_single_segment(base_draft, draft, target_index, aspect_ratio, effective_duration)
continue
raw, _provider_response, draft = routed.value
except Exception as exc: # noqa: BLE001
_fail_task(task, reservation, str(exc))
settled = True
@@ -819,7 +851,7 @@ def regenerate_segment_via_agent(*, project, user, model_config: ModelConfig, se
stream_script_agent target_index 分支同源,但同步返回(不走 SSE)计费 reservecharge/release 闭环"""
from django.db import transaction
from apps.ai.services import build_provider, create_ai_task
from apps.ai.services import create_ai_task, execute_routed_text_request
from apps.billing.services.ledger import charge_reserved_credit
base_draft = _draft_from_version(segment.script_version) # 用 DB 行重建基准,别用可能 stale 的 content
@@ -850,18 +882,40 @@ def regenerate_segment_via_agent(*, project, user, model_config: ModelConfig, se
"endpoint": model_config.endpoint,
"mode": "revise",
"target_index": target_index,
"model_routing_v1": True,
},
)
reservation = task.credit_reservation
# 每条真实尝试的平台成本由统一执行器累计;用户积分仍只结算这一条脚本任务。
task.base_cost = Decimal("0")
task.save(update_fields=["base_cost", "updated_at"])
# 实际平台成本由每条 AIModelAttempt 累加;用户积分仍只结算这一条逻辑任务。
task.base_cost = Decimal("0")
task.save(update_fields=["base_cost", "updated_at"])
try:
task.status = AITask.Status.SUBMITTED
task.submitted_at = timezone.now()
task.save(update_fields=["status", "submitted_at", "updated_at"])
provider = build_provider(model_config)
response = provider.chat_completion(model=model_config.name, endpoint=model_config.endpoint, messages=messages)
raw = provider.extract_text(response)
draft = normalize_draft(raw, aspect_ratio=aspect_ratio, total_duration=total_duration)
draft = _merge_single_segment(base_draft, draft, target_index, aspect_ratio, total_duration)
def validate_segment_text(raw_text: str) -> dict:
candidate = normalize_draft(raw_text, aspect_ratio=aspect_ratio, total_duration=total_duration)
return _merge_single_segment(base_draft, candidate, target_index, aspect_ratio, total_duration)
routed = execute_routed_text_request(
task=task,
primary_model=model_config,
messages=messages,
streaming=False,
structured_output=True,
business_operation="script_generate",
temperature=0.3,
validate_text=validate_segment_text,
request_summary={
"mode": "revise",
"target_index": target_index,
"base_version_id": str(segment.script_version_id),
},
)
raw, _response, draft = routed.value
with transaction.atomic():
task.status = AITask.Status.SUCCEEDED
task.response_payload = {"raw": raw[:8000]}
@@ -871,6 +925,6 @@ def regenerate_segment_via_agent(*, project, user, model_config: ModelConfig, se
charge_reserved_credit(reservation=reservation, actual_amount=task.actual_cost)
script = persist_script_draft(project=project, user=user, task=task, draft=draft, source="revise")
return script
except Exception:
_fail_task(task, reservation, "单镜重跑失败")
except Exception as exc:
_fail_task(task, reservation, str(exc) or "单镜重跑失败")
raise
File diff suppressed because it is too large Load Diff
+4 -2
View File
@@ -63,7 +63,7 @@ def generate_model_triview_task(self, task_id: str) -> str:
@app.task(bind=True, max_retries=0)
def poll_free_video_task(self, task_id: str, attempt: int = 0) -> str:
"""自由创作视频·worker 兜底轮询:每 30s 一次自重排(不依赖 celery beat),
上限 60 (30 分钟,足够 Seedance 5-10 分钟出片)finalize 幂等(POSTPROCESSING 认领),
总等待上限由统一视频 generation_timeout 配置控制finalize 幂等(POSTPROCESSING 认领),
与前端主动 poll 并存不双扣轮询本身出错不重试(max_retries=0),下一次自重排继续"""
from apps.ai.free_video import finalize_free_video
from apps.ai.models import AITask
@@ -82,7 +82,9 @@ def poll_free_video_task(self, task_id: str, attempt: int = 0) -> str:
if getattr(dj_settings, "CELERY_TASK_ALWAYS_EAGER", False):
return task_id
if task.status in (AITask.Status.SUBMITTED, AITask.Status.POLLING) and attempt < 60:
# 是否终止由 finalize_free_video 读取统一 generation_timeout 决定;这里不再维护一套
# “60 次轮询”硬编码,避免修改成片等待时限时 Worker 与 Web 轮询口径不一致。
if task.status in (AITask.Status.SUBMITTED, AITask.Status.POLLING):
poll_free_video_task.apply_async(args=[task_id, attempt + 1], countdown=30)
return task_id
@@ -0,0 +1,248 @@
from decimal import Decimal
from unittest.mock import Mock, patch
import requests
from django.test import TestCase
from apps.accounts.models import Team, TeamMember, User
from apps.ai.models import AITask, ModelConfig, ModelProvider
from apps.ai.services import generate_base_asset, run_base_asset_task
from apps.assets.models import Asset, AssetFile
from apps.billing.models import CreditAccount, CreditLedger
from apps.products.models import Product
from apps.projects.models import BaseAssetGroup, Project
def _image_metadata(*, outbound=True, priority=20, base_cost="0.50"):
return {
"routing": {
"fallback_on_failure": outbound,
"fallback_candidate": True,
},
"capabilities": {
"operations": ["image_generate", "image_edit"],
"features": [],
"reference_modes": ["none", "single", "multiple"],
"max_reference_images": 9,
"aspect_ratios": ["1:1", "3:4", "4:5", "9:16", "16:9"],
},
"pricing": {"base_cost_yuan": base_cost},
"test_priority": priority,
}
class BaseAssetRoutingTests(TestCase):
def setUp(self):
ModelConfig.objects.filter(capability=ModelConfig.Capability.IMAGE).update(
status=ModelConfig.Status.DISABLED
)
self.user = User.objects.create_user(username="base-routing", password="x")
self.team = Team.objects.create(name="Base Routing", owner=self.user)
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
CreditAccount.objects.create(team=self.team, balance=Decimal("1000"))
self.product = Product.objects.create(
team=self.team,
created_by=self.user,
title="橙色童装上衣",
category="服饰内衣",
)
self.project = Project.objects.create(
team=self.team,
created_by=self.user,
product=self.product,
name="基础资产路由项目",
)
self.provider_mocks = {}
patch("apps.ai.services.get_image_provider", side_effect=self._provider_for).start()
patch("apps.ai.services._store_generated_media", side_effect=self._store_media).start()
patch("apps.ai.tasks.generate_base_asset_task.delay").start()
self.addCleanup(patch.stopall)
def provider(self, name, priority):
return ModelProvider.objects.create(
name=name,
display_name=name,
status=ModelProvider.Status.ACTIVE,
metadata={"routing": {"fallback_priority": priority}},
)
def model(self, provider, name, *, outbound=True, base_cost="0.50"):
return ModelConfig.objects.create(
provider=provider,
name=name,
display_name=name,
capability=ModelConfig.Capability.IMAGE,
endpoint="images/generations",
unit_price=Decimal("20"),
status=ModelConfig.Status.ACTIVE,
metadata=_image_metadata(
outbound=outbound,
priority=(provider.metadata or {}).get("routing", {}).get("fallback_priority", 20),
base_cost=base_cost,
),
)
@staticmethod
def _new_provider_mock():
provider = Mock()
provider.image_generation.return_value = {"data": [{"url": "http://example.test/out.png"}]}
provider.image_edit.return_value = {"data": [{"url": "http://example.test/out.png"}]}
provider.extract_first_media_url.side_effect = lambda response: response["data"][0]["url"]
return provider
@staticmethod
def _new_generation_only_provider_mock():
provider = Mock(spec=["image_generation", "extract_first_media_url"])
provider.image_generation.return_value = {"data": [{"url": "http://example.test/out.png"}]}
provider.extract_first_media_url.side_effect = lambda response: response["data"][0]["url"]
return provider
def _provider_for(self, model):
return self.provider_mocks.setdefault(model.id, self._new_provider_mock())
def _store_media(self, **kwargs):
return Asset.objects.create(
team=kwargs["team"],
created_by=kwargs["user"],
name=kwargs["name"],
asset_type=kwargs["asset_type"],
source=Asset.Source.AI_GENERATED,
category=kwargs["category"],
origin_task=kwargs["task"],
)
def reference_asset(self, name="商品主图", category=Asset.Category.PRODUCT_IMAGE):
asset = Asset.objects.create(
team=self.team,
created_by=self.user,
name=name,
asset_type=Asset.Type.IMAGE,
source=Asset.Source.UPLOAD,
category=category,
)
AssetFile.objects.create(
asset=asset,
object_key=f"{name}.png",
bucket="bucket",
content_type="image/png",
preview_url=f"http://example.test/{name}.png",
is_primary=True,
)
return asset
def submit(self, kind, *, reference_asset_id=None):
return generate_base_asset(
project=self.project,
user=self.user,
kind=kind,
prompt=f"{kind} 基础资产提示词",
label=f"{kind}-label",
reference_asset_id=reference_asset_id,
)
@staticmethod
def ledger_count(task, ledger_type):
return CreditLedger.objects.filter(task=task, ledger_type=ledger_type).count()
def test_product_reference_uses_routed_image_edit_and_keeps_group(self):
primary = self.model(self.provider("base-product-primary", 20), "base-product")
cover = self.reference_asset()
self.product.cover_asset = cover
self.product.save(update_fields=["cover_asset"])
task = self.submit(BaseAssetGroup.Kind.PRODUCT)
run_base_asset_task(task_id=str(task.id))
task.refresh_from_db()
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
self.assertTrue(task.request_payload["model_routing_v1"])
attempt = task.model_attempts.get()
self.assertEqual(attempt.model_config_id, primary.id)
self.assertEqual(attempt.operation, "image_edit")
self.assertEqual(attempt.request_summary["reference_images"], 1)
self.assertEqual(attempt.request_summary["base_asset_kind"], "product")
call = self.provider_mocks[primary.id].image_edit.call_args
self.assertEqual(call.kwargs["images"], ["http://example.test/商品主图.png"])
group = BaseAssetGroup.objects.get(project=self.project, kind=BaseAssetGroup.Kind.PRODUCT)
self.assertEqual(group.adopted_asset.origin_task_id, task.id)
self.assertEqual(group.metadata["label"], "product-label")
def test_person_and_scene_without_reference_use_routed_generation(self):
primary = self.model(self.provider("base-generate-primary", 20), "base-generate")
for kind in (BaseAssetGroup.Kind.PERSON, BaseAssetGroup.Kind.SCENE):
with self.subTest(kind=kind):
task = self.submit(kind)
run_base_asset_task(task_id=str(task.id))
task.refresh_from_db()
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
attempt = task.model_attempts.get()
self.assertEqual(attempt.model_config_id, primary.id)
self.assertEqual(attempt.operation, "image_generate")
self.assertEqual(attempt.request_summary["reference_images"], 0)
self.assertEqual(attempt.request_summary["base_asset_kind"], kind)
group = BaseAssetGroup.objects.get(project=self.project, kind=kind)
self.assertEqual(group.adopted_asset.origin_task_id, task.id)
def test_person_rerun_reference_is_preserved_across_fallback(self):
primary = self.model(
self.provider("base-person-primary", 100), "person-primary", base_cost="0.25"
)
candidate = self.model(
self.provider("volcano", 10), "person-candidate", outbound=False, base_cost="0.75"
)
portrait = self.reference_asset("人物立绘", Asset.Category.PERSON)
task = self.submit(BaseAssetGroup.Kind.PERSON, reference_asset_id=str(portrait.id))
self.provider_mocks[primary.id] = self._new_provider_mock()
self.provider_mocks[primary.id].image_edit.side_effect = requests.ConnectionError("offline")
run_base_asset_task(task_id=str(task.id))
task.refresh_from_db()
attempts = list(task.model_attempts.all())
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
self.assertEqual([a.model_config_id for a in attempts], [primary.id, primary.id, candidate.id])
self.assertTrue(attempts[-1].is_fallback)
self.assertTrue(all(a.request_summary["reference_images"] == 1 for a in attempts))
self.assertTrue(all(a.public_model_name == "AirShelf Image" for a in attempts))
candidate_call = self.provider_mocks[candidate.id].image_edit.call_args
self.assertEqual(candidate_call.kwargs["images"], ["http://example.test/人物立绘.png"])
self.assertEqual(task.base_cost, Decimal("0.7500"))
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
def test_all_base_asset_candidates_fail_release_once(self):
primary = self.model(self.provider("base-all-primary", 100), "base-primary")
candidate = self.model(self.provider("volcano", 10), "base-candidate", outbound=False)
task = self.submit(BaseAssetGroup.Kind.SCENE)
for model in (primary, candidate):
self.provider_mocks[model.id] = self._new_provider_mock()
self.provider_mocks[model.id].image_generation.side_effect = requests.ConnectionError("offline")
run_base_asset_task(task_id=str(task.id))
task.refresh_from_db()
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(task.model_attempts.count(), 3)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
def test_direct_product_primary_uses_image_input_and_public_real_name(self):
primary = self.model(
self.provider("volcano", 10), "seedream-base-product", outbound=False
)
self.provider_mocks[primary.id] = self._new_generation_only_provider_mock()
cover = self.reference_asset()
self.product.cover_asset = cover
self.product.save(update_fields=["cover_asset"])
task = self.submit(BaseAssetGroup.Kind.PRODUCT)
run_base_asset_task(task_id=str(task.id))
task.refresh_from_db()
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
call = self.provider_mocks[primary.id].image_generation.call_args
self.assertEqual(call.kwargs["image"], ["http://example.test/商品主图.png"])
self.assertEqual(task.model_attempts.get().public_model_name, primary.display_name)
@@ -0,0 +1,238 @@
from decimal import Decimal
from unittest.mock import Mock, patch
import requests
from django.test import TestCase
from apps.accounts.models import Team, TeamMember, User
from apps.ai.models import AITask, ModelConfig, ModelProvider
from apps.ai.services import EXTRACT_TEXT_MODEL_NAME, run_extract_entities_task, submit_extract_entities
from apps.billing.models import CreditAccount, CreditLedger
from apps.products.models import Product
from apps.projects.models import Project, ScriptSegment, ScriptVersion
def _metadata(*, outbound=True, base_cost="0.50"):
return {
"routing": {"fallback_on_failure": outbound, "fallback_candidate": True},
"capabilities": {
"operations": ["chat"],
"features": ["streaming", "structured_output"],
},
"pricing": {"base_cost_yuan": base_cost},
}
class EntityExtractionRoutingTests(TestCase):
def setUp(self):
ModelConfig.objects.filter(capability=ModelConfig.Capability.TEXT).update(
status=ModelConfig.Status.DISABLED
)
self.user = User.objects.create_user(username="entity-routing", password="x")
self.team = Team.objects.create(name="Entity Routing", owner=self.user)
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
CreditAccount.objects.create(team=self.team, balance=Decimal("1000"))
product = Product.objects.create(team=self.team, created_by=self.user, title="测试商品")
self.project = Project.objects.create(
team=self.team,
created_by=self.user,
product=product,
name="实体提取路由项目",
metadata={"cast": ["旧角色"], "entities_extracted": False},
)
self.script = ScriptVersion.objects.create(
project=self.project,
title="脚本",
content="结构化脚本",
is_adopted=True,
)
self.segment = ScriptSegment.objects.create(
script_version=self.script,
sort_order=0,
narration="女主在客厅展示商品",
visual_prompt="女主站在客厅",
entity_refs=["old"],
)
self.provider_mocks = {}
patch("apps.ai.services.get_text_provider", side_effect=self._provider_for).start()
patch("apps.ai.tasks.extract_entities_task.delay").start()
self.addCleanup(patch.stopall)
def provider(self, name, priority):
return ModelProvider.objects.create(
name=name,
display_name=name,
status=ModelProvider.Status.ACTIVE,
metadata={"routing": {"fallback_priority": priority}},
)
def model(self, provider, name, *, outbound=True, base_cost="0.50"):
return ModelConfig.objects.create(
provider=provider,
name=name,
display_name=name,
capability=ModelConfig.Capability.TEXT,
endpoint="chat/completions",
unit_price=Decimal("10"),
status=ModelConfig.Status.ACTIVE,
metadata=_metadata(outbound=outbound, base_cost=base_cost),
)
@staticmethod
def valid_events():
return [
{
"type": "delta",
"text": (
'{"entities":['
'{"id":"c1","type":"character","name":"女主","visual_prompt":"都市女主"},'
'{"id":"s1","type":"scene","name":"客厅","visual_prompt":"现代客厅"}'
'],"segments":[{"index":0,"entity_refs":["c1","s1"]}]}'
),
},
{"type": "done"},
]
@staticmethod
def invalid_events():
return [{"type": "delta", "text": "这是一段没有 JSON 的解释"}, {"type": "done"}]
@classmethod
def _new_provider_mock(cls):
provider = Mock()
provider.chat_completion_stream.return_value = cls.valid_events()
return provider
def _provider_for(self, model):
return self.provider_mocks.setdefault(model.id, self._new_provider_mock())
def submit(self):
return submit_extract_entities(project=self.project, user=self.user)
@staticmethod
def ledger_count(task, ledger_type):
return CreditLedger.objects.filter(task=task, ledger_type=ledger_type).count()
def test_first_success_records_streaming_structured_attempt_and_persists_entities(self):
primary = self.model(self.provider("entity-primary", 20), "entity-primary")
task = self.submit()
run_extract_entities_task(task_id=str(task.id))
task.refresh_from_db()
self.project.refresh_from_db()
self.segment.refresh_from_db()
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
self.assertTrue(task.request_payload["model_routing_v1"])
attempt = task.model_attempts.get()
self.assertEqual(attempt.model_config_id, primary.id)
self.assertEqual(attempt.operation, "chat")
self.assertTrue(attempt.request_summary["streaming"])
self.assertTrue(attempt.request_summary["structured_output"])
self.assertEqual(attempt.request_summary["business_operation"], "entity_extract")
self.assertEqual(attempt.public_model_name, "AirShelf Script")
self.assertTrue(self.project.metadata["entities_extracted"])
self.assertEqual(self.project.metadata["cast"], ["女主"])
self.assertEqual(self.project.metadata["scenes"], ["客厅"])
self.assertEqual(self.segment.entity_refs, ["c1", "s1"])
call = self.provider_mocks[primary.id].chat_completion_stream.call_args
self.assertGreater(call.kwargs["timeout"], 0)
self.assertEqual(call.kwargs["temperature"], 0.3)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
def test_invalid_structure_retries_then_falls_back_and_records_failed_call_costs(self):
primary = self.model(
self.provider("entity-fallback-primary", 100),
"entity-primary",
base_cost="0.25",
)
candidate = self.model(
self.provider("entity-fallback-candidate", 10),
"entity-candidate",
outbound=False,
base_cost="0.75",
)
primary_mock = self._new_provider_mock()
primary_mock.chat_completion_stream.side_effect = [
self.invalid_events(),
self.invalid_events(),
self.invalid_events(),
]
self.provider_mocks[primary.id] = primary_mock
task = self.submit()
run_extract_entities_task(task_id=str(task.id))
task.refresh_from_db()
attempts = list(task.model_attempts.all())
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
self.assertEqual(
[a.model_config_id for a in attempts],
[primary.id, primary.id, primary.id, candidate.id],
)
self.assertEqual(
[a.status for a in attempts],
["failed", "failed", "failed", "succeeded"],
)
self.assertTrue(attempts[-1].is_fallback)
self.assertTrue(attempts[0].response_summary["validation_failed"])
self.assertEqual(task.base_cost, Decimal("1.5000"))
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
def test_all_candidates_fail_releases_and_preserves_old_project_metadata(self):
primary = self.model(self.provider("entity-all-primary", 100), "entity-primary")
fallback_1 = self.model(
self.provider("entity-all-fallback-1", 10), "entity-fallback-1", outbound=False
)
fallback_2 = self.model(
self.provider("entity-all-fallback-2", 20), "entity-fallback-2", outbound=False
)
for model in (primary, fallback_1, fallback_2):
provider = self._new_provider_mock()
provider.chat_completion_stream.side_effect = requests.ConnectionError("offline")
self.provider_mocks[model.id] = provider
task = self.submit()
run_extract_entities_task(task_id=str(task.id))
task.refresh_from_db()
self.project.refresh_from_db()
self.segment.refresh_from_db()
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(task.model_attempts.count(), 4)
self.assertEqual(self.project.metadata["cast"], ["旧角色"])
self.assertFalse(self.project.metadata["entities_extracted"])
self.assertEqual(self.segment.entity_refs, ["old"])
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
def test_direct_doubao_primary_retries_but_does_not_switch_when_outbound_disabled(self):
primary = self.model(
self.provider("doubao", 10), EXTRACT_TEXT_MODEL_NAME, outbound=False
)
candidate = self.model(
self.provider("entity-unused-candidate", 20), "entity-unused", outbound=False
)
provider = self._new_provider_mock()
provider.chat_completion_stream.side_effect = requests.ConnectionError("offline")
self.provider_mocks[primary.id] = provider
task = self.submit()
run_extract_entities_task(task_id=str(task.id))
task.refresh_from_db()
attempts = list(task.model_attempts.all())
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(
[attempt.model_config_id for attempt in attempts],
[primary.id, primary.id, primary.id],
)
self.assertFalse(any(attempt.model_config_id == candidate.id for attempt in attempts))
self.assertTrue(all(attempt.public_model_name == primary.display_name for attempt in attempts))
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
@@ -0,0 +1,243 @@
from copy import deepcopy
from datetime import timedelta
from decimal import Decimal
from unittest.mock import Mock, patch
import requests
from django.conf import settings
from django.test import TestCase, override_settings
from django.utils import timezone
from apps.accounts.models import Team, TeamMember, User
from apps.ai.free_video import finalize_free_video, submit_free_video
from apps.ai.models import AITask, ModelConfig, ModelProvider
from apps.assets.models import Asset
from apps.billing.models import CreditAccount, CreditLedger
STANDARD = "doubao-seedance-2-0-260128"
def _fast_routing_policy():
policy = deepcopy(settings.MODEL_ROUTING_POLICY)
policy["jitter_ratio"] = 0
policy["video"]["submit_retry_delays"] = [0]
return policy
def _candidate_metadata(*, outbound=False, max_images=9, max_videos=3, max_audios=3):
return {
"routing": {
"fallback_on_failure": outbound,
"fallback_candidate": True,
},
"capabilities": {
"operations": ["video_generate"],
"features": ["generate_audio"],
"max_reference_images": max_images,
"max_reference_videos": max_videos,
"max_reference_audios": max_audios,
"aspect_ratios": ["16:9"],
"resolutions": ["480p"],
"durations": [4],
},
"pricing": {
"unit": "cny_per_million_tokens",
"default": {"no_ref_video": 23, "with_ref_video": 14},
},
}
@override_settings(MODEL_ROUTING_POLICY=_fast_routing_policy())
class FreeVideoRoutingTests(TestCase):
def setUp(self):
self.user = User.objects.create_user(username="free-video-routing", password="x")
self.team = Team.objects.create(name="Free Video Routing", owner=self.user)
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
CreditAccount.objects.create(team=self.team, balance=Decimal("10000"))
ModelConfig.objects.filter(capability=ModelConfig.Capability.VIDEO).update(
status=ModelConfig.Status.DISABLED,
is_default=False,
)
self.primary = ModelConfig.objects.select_related("provider").get(
name=STANDARD,
capability=ModelConfig.Capability.VIDEO,
)
primary_metadata = deepcopy(self.primary.metadata or {})
primary_metadata.setdefault("routing", {})["fallback_on_failure"] = True
primary_metadata["routing"]["fallback_candidate"] = True
self.primary.metadata = primary_metadata
self.primary.status = ModelConfig.Status.ACTIVE
self.primary.is_default = True
self.primary.provider.status = ModelProvider.Status.ACTIVE
self.primary.provider.metadata = {"routing": {"fallback_priority": 10}}
self.primary.provider.save(update_fields=["status", "metadata", "updated_at"])
self.primary.save(update_fields=["metadata", "status", "is_default", "updated_at"])
self.provider_mocks = {}
patch("apps.ai.services.get_video_provider", side_effect=self._provider_for).start()
patch("apps.ai.tasks.poll_free_video_task.apply_async").start()
patch("apps.ai.free_video._notify_failure").start()
self.addCleanup(patch.stopall)
def candidate(self, name="free-video-candidate", **metadata_overrides):
provider = ModelProvider.objects.create(
name=f"{name}-provider",
display_name=name,
status=ModelProvider.Status.ACTIVE,
base_url="https://video.example/v1",
metadata={"routing": {"fallback_priority": 20}},
)
metadata = _candidate_metadata(**metadata_overrides)
return ModelConfig.objects.create(
provider=provider,
name=name,
display_name=name,
capability=ModelConfig.Capability.VIDEO,
endpoint="videos",
status=ModelConfig.Status.ACTIVE,
metadata=metadata,
)
def _provider_for(self, model):
return self.provider_mocks.setdefault(model.id, Mock())
def params(self, **overrides):
params = {
"prompt": "A product video",
"mode": "universal",
"model": STANDARD,
"aspect_ratio": "16:9",
"resolution": "480p",
"duration": 4,
"generate_audio": True,
"references": [],
}
params.update(overrides)
return params
@staticmethod
def ledger_count(task, ledger_type):
return CreditLedger.objects.filter(task=task, ledger_type=ledger_type).count()
def test_multi_material_fallback_uses_actual_model_for_poll_and_settlement(self):
candidate = self.candidate()
primary_provider = self._provider_for(self.primary)
fallback_provider = self._provider_for(candidate)
primary_provider.create_video_task.side_effect = requests.ConnectionError("primary offline")
fallback_provider.create_video_task.return_value = {"id": "free-fallback", "status": "queued"}
references = [
{"url": "https://media.example/ref.png", "type": "image", "label": "image"},
{"url": "https://media.example/ref.mp4", "type": "video", "label": "video", "duration": 1},
{"url": "https://media.example/ref.mp3", "type": "audio", "label": "audio"},
]
task = submit_free_video(team=self.team, user=self.user, params=self.params(references=references))
self.assertEqual(task.status, AITask.Status.SUBMITTED)
self.assertEqual(task.provider_task_id, "free-fallback")
self.assertEqual(task.request_payload["actual_model_config_id"], str(candidate.id))
attempts = list(task.model_attempts.order_by("sequence"))
self.assertEqual(
[attempt.model_config_id for attempt in attempts],
[self.primary.id, self.primary.id, candidate.id],
)
self.assertEqual(attempts[-1].request_summary["reference_images"], 1)
self.assertEqual(attempts[-1].request_summary["reference_videos"], 1)
self.assertEqual(attempts[-1].request_summary["reference_audios"], 1)
submit_call = fallback_provider.create_video_task.call_args.kwargs
self.assertEqual(len(submit_call["content_items"]), 3)
self.assertEqual(submit_call["timeout"], 120.0)
fallback_provider.poll_video_task.return_value = {
"status": "succeeded",
"usage": {"total_tokens": 30000},
"content": {"video_url": "https://video.example/result.mp4"},
}
fallback_provider.extract_first_media_url.return_value = "https://video.example/result.mp4"
with patch("apps.ai.free_video._store_free_video_media") as store_media:
store_media.return_value = Asset.objects.create(
team=self.team,
created_by=self.user,
name="free result",
asset_type=Asset.Type.VIDEO,
source=Asset.Source.AI_GENERATED,
category=Asset.Category.FREE_CREATE,
origin_task=task,
)
result = finalize_free_video(task=task)
self.assertEqual(result.status, AITask.Status.SUCCEEDED)
self.assertFalse(primary_provider.poll_video_task.called)
fallback_provider.poll_video_task.assert_called_once_with(
endpoint=candidate.endpoint,
provider_task_id="free-fallback",
timeout=60.0,
)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
def test_candidate_missing_audio_capacity_is_filtered_out(self):
candidate = self.candidate("free-video-incompatible", max_audios=0)
primary_provider = self._provider_for(self.primary)
unused_provider = self._provider_for(candidate)
primary_provider.create_video_task.side_effect = requests.ConnectionError("primary offline")
references = [
{"url": "https://media.example/ref.png", "type": "image"},
{"url": "https://media.example/ref.mp3", "type": "audio"},
]
task = submit_free_video(team=self.team, user=self.user, params=self.params(references=references))
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertFalse(unused_provider.create_video_task.called)
self.assertEqual(task.model_attempts.count(), 2)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
def test_submit_read_timeout_is_not_retried_or_fallbacked(self):
candidate = self.candidate("free-video-timeout-unused")
primary_provider = self._provider_for(self.primary)
unused_provider = self._provider_for(candidate)
primary_provider.create_video_task.side_effect = requests.ReadTimeout("response lost")
task = submit_free_video(team=self.team, user=self.user, params=self.params())
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(task.model_attempts.count(), 1)
self.assertFalse(unused_provider.create_video_task.called)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
def test_poll_network_error_keeps_inflight_task_and_reservation(self):
provider = self._provider_for(self.primary)
provider.create_video_task.return_value = {"id": "free-poll", "status": "queued"}
task = submit_free_video(team=self.team, user=self.user, params=self.params())
provider.poll_video_task.side_effect = requests.ConnectionError("poll offline")
with self.assertRaises(requests.ConnectionError):
finalize_free_video(task=task)
task.refresh_from_db()
self.assertEqual(task.status, AITask.Status.SUBMITTED)
self.assertEqual(task.model_attempts.count(), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
self.assertGreater(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
def test_generation_timeout_releases_without_polling_or_resubmitting(self):
provider = self._provider_for(self.primary)
provider.create_video_task.return_value = {"id": "free-slow", "status": "queued"}
task = submit_free_video(team=self.team, user=self.user, params=self.params())
task.submitted_at = timezone.now() - timedelta(seconds=1900)
task.save(update_fields=["submitted_at", "updated_at"])
result = finalize_free_video(task=task)
result.refresh_from_db()
self.assertEqual(result.status, AITask.Status.FAILED)
self.assertFalse(provider.poll_video_task.called)
self.assertEqual(provider.create_video_task.call_count, 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
+348
View File
@@ -0,0 +1,348 @@
from copy import deepcopy
from datetime import timedelta
from django.conf import settings
from django.test import SimpleTestCase, TestCase, override_settings
from django.utils import timezone
from apps.adminpanel.serializers import AdminModelConfigSerializer, AdminModelProviderSerializer
from apps.ai.catalog import VOLCANO_MODELS, VOLCANO_PROVIDER, YUNQI_MODELS, YUNQI_PROVIDER
from apps.ai.model_routing import (
ModelRequirements,
match_model_requirements,
model_metadata_errors,
provider_metadata_errors,
resolve_fallback_candidates,
)
from apps.ai.models import ModelConfig, ModelProvider
def text_metadata(*, outbound=True, candidate=True):
return {
"routing": {
"fallback_on_failure": outbound,
"fallback_candidate": candidate,
},
"capabilities": {
"operations": ["chat"],
"features": ["streaming", "structured_output"],
},
}
class ModelRequirementsMatchTests(SimpleTestCase):
def model(self, capability, metadata):
return ModelConfig(name="match-model", display_name="Match", capability=capability, metadata=metadata)
def test_requirements_reject_invalid_reference_and_limit_values(self):
with self.assertRaisesRegex(ValueError, "reference_images 不能为负数"):
ModelRequirements(capability="image", operation="image_generate", reference_images=-1)
with self.assertRaisesRegex(ValueError, "reference_images 必须等于 1"):
ModelRequirements(
capability="image",
operation="image_edit",
reference_mode="single",
reference_images=2,
)
def test_image_requires_exact_operation_reference_mode_count_and_ratio(self):
model = self.model(
ModelConfig.Capability.IMAGE,
{
"capabilities": {
"operations": ["image_generate", "image_edit"],
"features": [],
"reference_modes": ["none", "single", "multiple"],
"max_reference_images": 4,
"aspect_ratios": ["1:1", "9:16"],
}
},
)
matched = match_model_requirements(
model,
ModelRequirements(
capability="image",
operation="image_edit",
reference_mode="multiple",
reference_images=4,
aspect_ratio="9:16",
),
)
rejected = match_model_requirements(
model,
ModelRequirements(
capability="image",
operation="image_edit",
reference_mode="multiple",
reference_images=5,
aspect_ratio="16:9",
),
)
self.assertTrue(matched.matched)
self.assertFalse(rejected.matched)
self.assertTrue(any("参考图片上限不足" in reason for reason in rejected.reasons))
self.assertTrue(any("画面比例" in reason for reason in rejected.reasons))
def test_video_requires_features_resolution_duration_and_exact_price_tier(self):
model = self.model(
ModelConfig.Capability.VIDEO,
{
"pricing": {"default": {"no_ref_video": 1}, "1080p": {"no_ref_video": 2}},
"capabilities": {
"operations": ["video_generate"],
"features": ["text_to_video", "start_frame", "generate_audio"],
"resolutions": ["720p", "1080p", "4k"],
"durations": [5, 10],
"aspect_ratios": ["16:9"],
},
},
)
matched = match_model_requirements(
model,
ModelRequirements(
capability="video",
operation="video_generate",
features=frozenset({"start_frame", "generate_audio"}),
resolution="1080p",
duration=10,
aspect_ratio="16:9",
),
)
rejected = match_model_requirements(
model,
ModelRequirements(
capability="video",
operation="video_generate",
features=frozenset({"video_reference"}),
resolution="4k",
duration=10,
aspect_ratio="16:9",
),
)
self.assertTrue(matched.matched)
self.assertFalse(rejected.matched)
self.assertTrue(any("video_reference" in reason for reason in rejected.reasons))
self.assertTrue(any("4k 的视频价格" in reason for reason in rejected.reasons))
def test_audio_requires_public_voice_mapping_and_limits(self):
model = self.model(
ModelConfig.Capability.AUDIO,
{
"capabilities": {
"operations": ["tts"],
"features": [],
"languages": ["zh-CN"],
"voice_map": {"narrator": "provider-voice-7"},
"max_chars": 1000,
"speed_range": [0.5, 2.0],
"output_formats": ["mp3"],
}
},
)
matched = match_model_requirements(
model,
ModelRequirements(
capability="audio",
operation="tts",
language="zh-CN",
public_voice="narrator",
char_count=999,
speed_ratio=1.5,
output_format="mp3",
),
)
rejected = match_model_requirements(
model,
ModelRequirements(
capability="audio",
operation="tts",
language="zh-CN",
public_voice="unknown",
char_count=1001,
speed_ratio=3,
output_format="wav",
),
)
self.assertTrue(matched.matched)
self.assertFalse(rejected.matched)
self.assertGreaterEqual(len(rejected.reasons), 4)
def test_missing_capability_contract_is_conservatively_rejected(self):
model = self.model(ModelConfig.Capability.TEXT, {"routing": {"fallback_candidate": True}})
result = match_model_requirements(model, ModelRequirements(capability="text", operation="chat"))
self.assertFalse(result.matched)
self.assertIn("不支持操作 chat", result.reasons)
def test_chinese_metadata_validation_catches_bad_switches_and_priority(self):
self.assertTrue(provider_metadata_errors({"routing": {"fallback_priority": True}}))
errors = model_metadata_errors(
"video",
{
"routing": {"fallback_on_failure": "yes", "fallback_candidate": True},
"capabilities": {"operations": ["video_generate"]},
},
)
self.assertTrue(any("true 或 false" in error for error in errors))
self.assertTrue(any("resolutions" in error for error in errors))
self.assertTrue(any("durations" in error for error in errors))
def test_bootstrap_catalog_contains_persistent_routing_contracts(self):
self.assertEqual(VOLCANO_PROVIDER["metadata"]["routing"]["fallback_priority"], 10)
self.assertEqual(YUNQI_PROVIDER["metadata"]["routing"]["fallback_priority"], 20)
for item in VOLCANO_MODELS + YUNQI_MODELS:
self.assertIn("routing", item["metadata"])
self.assertIn("capabilities", item["metadata"])
self.assertTrue(item["metadata"]["routing"]["fallback_candidate"])
self.assertTrue(item["metadata"]["capabilities"]["operations"])
class DynamicFallbackCandidateTests(TestCase):
requirements = ModelRequirements(capability="text", operation="chat")
def setUp(self):
# 数据迁移会种入真实候选;路由单测只比较本用例创建的数据,避免环境顺序干扰。
ModelConfig.objects.all().update(status=ModelConfig.Status.DISABLED)
def provider(self, name, priority=100, status=ModelProvider.Status.ACTIVE):
return ModelProvider.objects.create(
name=name,
display_name=name,
status=status,
metadata={"routing": {"fallback_priority": priority}},
)
def model(self, provider, name, *, outbound=True, candidate=True, status=ModelConfig.Status.ACTIVE, metadata=None):
return ModelConfig.objects.create(
provider=provider,
name=name,
display_name=name,
capability=ModelConfig.Capability.TEXT,
endpoint="chat/completions",
status=status,
metadata=metadata if metadata is not None else text_metadata(outbound=outbound, candidate=candidate),
)
def policy(self, max_models=5):
raw = deepcopy(settings.MODEL_ROUTING_POLICY)
raw["max_models"] = max_models
raw["max_calls"] = max(raw["max_calls"], max_models)
return raw
def test_dynamic_order_is_priority_then_latest_updated_at(self):
primary = self.model(self.provider("route-primary", 500), "primary")
priority_10 = self.model(self.provider("route-priority-10", 10), "priority-10")
priority_20 = self.model(self.provider("route-priority-20", 20), "priority-20")
normal_provider = self.provider("route-normal", 100)
normal_old = self.model(normal_provider, "normal-old")
normal_new = self.model(normal_provider, "normal-new")
now = timezone.now()
ModelConfig.objects.filter(pk=normal_old.pk).update(updated_at=now - timedelta(days=1))
ModelConfig.objects.filter(pk=normal_new.pk).update(updated_at=now)
with override_settings(MODEL_ROUTING_POLICY=self.policy()):
candidates = resolve_fallback_candidates(primary_model=primary, requirements=self.requirements)
self.assertEqual(
[model.name for model in candidates],
[priority_10.name, priority_20.name, normal_new.name, normal_old.name],
)
def test_user_primary_is_preserved_and_outbound_switch_can_stop_fallback(self):
primary = self.model(self.provider("route-stop-primary", 500), "primary", outbound=False)
self.model(self.provider("route-stop-candidate", 1), "candidate")
candidates = resolve_fallback_candidates(primary_model=primary, requirements=self.requirements)
self.assertEqual(candidates, [])
self.assertEqual(primary.name, "primary")
def test_attempted_models_excluded_providers_and_disabled_rows_are_filtered(self):
primary = self.model(self.provider("route-filter-primary", 500), "primary")
attempted = self.model(self.provider("route-filter-attempted", 1), "attempted")
excluded = self.model(self.provider("route-filter-provider", 2), "excluded")
disabled_model = self.model(
self.provider("route-filter-disabled-model", 3),
"disabled-model",
status=ModelConfig.Status.DISABLED,
)
disabled_provider = self.model(
self.provider("route-filter-disabled-provider", 4, ModelProvider.Status.DISABLED),
"disabled-provider",
)
valid = self.model(self.provider("route-filter-valid", 5), "valid")
with override_settings(MODEL_ROUTING_POLICY=self.policy()):
candidates = resolve_fallback_candidates(
primary_model=primary,
requirements=self.requirements,
attempted_model_ids={attempted.id},
excluded_provider_ids={excluded.provider_id},
)
self.assertEqual([model.id for model in candidates], [valid.id])
self.assertNotIn(disabled_model.id, [model.id for model in candidates])
self.assertNotIn(disabled_provider.id, [model.id for model in candidates])
def test_new_openai_compatible_provider_enters_pool_by_data_only(self):
primary = self.model(self.provider("route-data-primary", 500), "primary")
future_provider = self.provider("future-openai-compatible-provider", 15)
future_model = self.model(future_provider, "future-text-model")
candidates = resolve_fallback_candidates(primary_model=primary, requirements=self.requirements)
self.assertEqual([model.id for model in candidates], [future_model.id])
def test_unconfigured_or_incompatible_candidates_are_conservatively_excluded(self):
primary = self.model(self.provider("route-safe-primary", 500), "primary")
self.model(self.provider("route-safe-missing", 1), "missing", metadata={})
self.model(
self.provider("route-safe-incompatible", 2),
"incompatible",
metadata={
"routing": {"fallback_on_failure": True, "fallback_candidate": True},
"capabilities": {"operations": ["embeddings"], "features": []},
},
)
candidates = resolve_fallback_candidates(primary_model=primary, requirements=self.requirements)
self.assertEqual(candidates, [])
def test_admin_serializers_reject_invalid_metadata_with_chinese_message(self):
provider = self.provider("route-admin", 100)
model = self.model(provider, "admin-model")
provider_serializer = AdminModelProviderSerializer(
provider,
data={"metadata": {"routing": {"fallback_priority": "first"}}},
partial=True,
)
model_serializer = AdminModelConfigSerializer(
model,
data={
"metadata": {
"routing": {"fallback_candidate": True, "fallback_on_failure": True},
"capabilities": {},
}
},
partial=True,
)
self.assertFalse(provider_serializer.is_valid())
self.assertIn("数字越小越优先", str(provider_serializer.errors))
self.assertFalse(model_serializer.is_valid())
self.assertIn("至少配置一个操作", str(model_serializer.errors))
def test_data_migration_seeded_existing_yunqi_rows(self):
provider = ModelProvider.objects.filter(name="yunqi").first()
model = ModelConfig.objects.filter(provider__name="yunqi", capability="image").first()
self.assertIsNotNone(provider)
self.assertEqual(provider.metadata["routing"]["fallback_priority"], 20)
self.assertIsNotNone(model)
self.assertIn("routing", model.metadata)
self.assertIn("capabilities", model.metadata)
@@ -0,0 +1,216 @@
from decimal import Decimal
from io import BytesIO
from unittest.mock import Mock, patch
import requests
from django.test import TestCase
from apps.accounts.models import Team, TeamMember, User
from apps.ai.models import AITask, ModelConfig, ModelProvider
from apps.ai.services import generate_model_triview, run_model_triview_task
from apps.assets.models import Asset, AssetFile, Model
from apps.billing.models import CreditAccount, CreditLedger
def _metadata(*, outbound=True, base_cost="0.50"):
return {
"routing": {"fallback_on_failure": outbound, "fallback_candidate": True},
"capabilities": {
"operations": ["image_generate", "image_edit"],
"features": [],
"reference_modes": ["none", "single", "multiple"],
"max_reference_images": 9,
"aspect_ratios": ["1:1", "3:4", "4:5", "9:16", "16:9"],
},
"pricing": {"base_cost_yuan": base_cost},
}
class ModelTriviewRoutingTests(TestCase):
def setUp(self):
ModelConfig.objects.filter(capability=ModelConfig.Capability.IMAGE).update(
status=ModelConfig.Status.DISABLED
)
self.user = User.objects.create_user(username="tri-routing", password="x")
self.team = Team.objects.create(name="Tri Routing", owner=self.user)
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
CreditAccount.objects.create(team=self.team, balance=Decimal("1000"))
self.portrait = Asset.objects.create(
team=self.team,
created_by=self.user,
name="模特形象",
asset_type=Asset.Type.IMAGE,
source=Asset.Source.UPLOAD,
category=Asset.Category.MODEL_PORTRAIT,
in_library=False,
)
AssetFile.objects.create(
asset=self.portrait,
object_key="portrait.png",
bucket="bucket",
content_type="image/png",
preview_url="http://example.test/portrait.png",
is_primary=True,
)
self.model_entity = Model.objects.create(
team=self.team,
created_by=self.user,
name="路由模特",
portrait_asset=self.portrait,
)
self.provider_mocks = {}
patch("apps.ai.services.get_image_provider", side_effect=self._provider_for).start()
patch("apps.ai.tasks.generate_model_triview_task.delay").start()
patch(
"apps.ai.services.VolcanoArkProvider.media_to_bytes",
return_value=(BytesIO(b"image"), "image/png"),
).start()
storage = patch("apps.ai.services.TosStorage").start()
stored = storage.return_value.upload_fileobj.return_value
stored.object_key = "tri.png"
stored.bucket = "bucket"
stored.content_type = "image/png"
stored.size_bytes = 5
patch("apps.assets.review.submit_asset_for_review").start()
self.addCleanup(patch.stopall)
def provider(self, name, priority):
return ModelProvider.objects.create(
name=name,
display_name=name,
status=ModelProvider.Status.ACTIVE,
metadata={"routing": {"fallback_priority": priority}},
)
def model(self, provider, name, *, outbound=True, base_cost="0.50"):
return ModelConfig.objects.create(
provider=provider,
name=name,
display_name=name,
capability=ModelConfig.Capability.IMAGE,
endpoint="images/generations",
unit_price=Decimal("20"),
status=ModelConfig.Status.ACTIVE,
metadata=_metadata(outbound=outbound, base_cost=base_cost),
)
@staticmethod
def _new_provider_mock():
provider = Mock()
provider.image_edit.return_value = {"data": [{"url": "http://example.test/tri.png"}]}
provider.image_generation.return_value = {"data": [{"url": "http://example.test/tri.png"}]}
provider.extract_first_media_url.side_effect = lambda response: response["data"][0]["url"]
return provider
@staticmethod
def _new_generation_only_provider_mock():
provider = Mock(spec=["image_generation", "extract_first_media_url"])
provider.image_generation.return_value = {"data": [{"url": "http://example.test/tri.png"}]}
provider.extract_first_media_url.side_effect = lambda response: response["data"][0]["url"]
return provider
def _provider_for(self, model):
return self.provider_mocks.setdefault(model.id, self._new_provider_mock())
def submit(self):
task, created = generate_model_triview(model=self.model_entity, user=self.user)
self.assertTrue(created)
return task
@staticmethod
def ledger_count(task, ledger_type):
return CreditLedger.objects.filter(task=task, ledger_type=ledger_type).count()
def test_first_success_records_single_reference_16_9_and_updates_model(self):
primary = self.model(self.provider("tri-primary", 20), "tri-primary-model")
task = self.submit()
run_model_triview_task(task_id=str(task.id))
task.refresh_from_db()
self.model_entity.refresh_from_db()
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
self.assertTrue(task.request_payload["model_routing_v1"])
attempt = task.model_attempts.get()
self.assertEqual(attempt.model_config_id, primary.id)
self.assertEqual(attempt.operation, "image_edit")
self.assertEqual(attempt.request_summary["reference_images"], 1)
self.assertEqual(attempt.request_summary["aspect_ratio"], "16:9")
self.assertTrue(attempt.request_summary["model_triview"])
call = self.provider_mocks[primary.id].image_edit.call_args
self.assertEqual(call.kwargs["images"], ["http://example.test/portrait.png"])
self.assertEqual(call.kwargs["size"], "1536x864")
self.assertIsNotNone(self.model_entity.triview_asset_id)
self.assertEqual(self.model_entity.triview_asset.origin_task_id, task.id)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
def test_retry_then_dynamic_candidate_success_charges_once(self):
primary = self.model(self.provider("tri-fallback-primary", 100), "tri-primary", base_cost="0.25")
candidate = self.model(
self.provider("volcano", 10), "tri-candidate", outbound=False, base_cost="0.75"
)
task = self.submit()
self.provider_mocks[primary.id] = self._new_provider_mock()
self.provider_mocks[primary.id].image_edit.side_effect = requests.ConnectionError("offline")
run_model_triview_task(task_id=str(task.id))
task.refresh_from_db()
attempts = list(task.model_attempts.all())
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
self.assertEqual([a.model_config_id for a in attempts], [primary.id, primary.id, candidate.id])
self.assertTrue(attempts[-1].is_fallback)
self.assertTrue(all(a.request_summary["aspect_ratio"] == "16:9" for a in attempts))
self.assertTrue(all(a.public_model_name == "AirShelf Image" for a in attempts))
self.assertEqual(task.base_cost, Decimal("0.7500"))
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
def test_all_candidates_fail_releases_and_keeps_existing_triview(self):
primary = self.model(self.provider("tri-all-primary", 100), "tri-primary")
candidate = self.model(self.provider("volcano", 10), "tri-candidate", outbound=False)
old = Asset.objects.create(
team=self.team,
created_by=self.user,
name="旧三视图",
asset_type=Asset.Type.IMAGE,
source=Asset.Source.AI_GENERATED,
category=Asset.Category.TRI_VIEW,
in_library=False,
)
self.model_entity.triview_asset = old
self.model_entity.save(update_fields=["triview_asset"])
task = self.submit()
for model in (primary, candidate):
self.provider_mocks[model.id] = self._new_provider_mock()
self.provider_mocks[model.id].image_edit.side_effect = requests.ConnectionError("offline")
run_model_triview_task(task_id=str(task.id))
task.refresh_from_db()
self.model_entity.refresh_from_db()
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(task.model_attempts.count(), 3)
self.assertEqual(self.model_entity.triview_asset_id, old.id)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
def test_direct_primary_uses_image_generation_without_outbound_fallback(self):
primary = self.model(
self.provider("volcano", 10), "seedream-triview", outbound=False
)
self.provider_mocks[primary.id] = self._new_generation_only_provider_mock()
task = self.submit()
run_model_triview_task(task_id=str(task.id))
task.refresh_from_db()
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
call = self.provider_mocks[primary.id].image_generation.call_args
self.assertEqual(call.kwargs["image"], ["http://example.test/portrait.png"])
self.assertEqual(call.kwargs["size"], "1536x864")
self.assertEqual(task.model_attempts.get().public_model_name, primary.display_name)
@@ -0,0 +1,234 @@
from decimal import Decimal
from unittest.mock import Mock, patch
import requests
from django.test import TestCase
from apps.accounts.models import Team, TeamMember, User
from apps.ai.models import AITask, ModelConfig, ModelProvider
from apps.ai.services import generate_person_triview, run_triview_task
from apps.assets.models import Asset, AssetFile, Model
from apps.billing.models import CreditAccount, CreditLedger
from apps.products.models import Product
from apps.projects.models import BaseAssetGroup, Project
def _metadata(*, outbound=True, base_cost="0.50"):
return {
"routing": {"fallback_on_failure": outbound, "fallback_candidate": True},
"capabilities": {
"operations": ["image_generate", "image_edit"],
"features": [],
"reference_modes": ["none", "single", "multiple"],
"max_reference_images": 9,
"aspect_ratios": ["1:1", "3:4", "4:5", "9:16", "16:9"],
},
"pricing": {"base_cost_yuan": base_cost},
}
class ProjectPersonTriviewRoutingTests(TestCase):
def setUp(self):
ModelConfig.objects.filter(capability=ModelConfig.Capability.IMAGE).update(
status=ModelConfig.Status.DISABLED
)
self.user = User.objects.create_user(username="project-tri-routing", password="x")
self.team = Team.objects.create(name="Project Tri Routing", owner=self.user)
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
CreditAccount.objects.create(team=self.team, balance=Decimal("1000"))
product = Product.objects.create(team=self.team, created_by=self.user, title="测试商品")
self.project = Project.objects.create(
team=self.team,
created_by=self.user,
product=product,
name="项目角色三视图",
)
self.portrait = Asset.objects.create(
team=self.team,
created_by=self.user,
name="角色立绘",
asset_type=Asset.Type.IMAGE,
source=Asset.Source.AI_GENERATED,
category=Asset.Category.PERSON,
)
AssetFile.objects.create(
asset=self.portrait,
object_key="portrait.png",
bucket="bucket",
content_type="image/png",
preview_url="http://example.test/portrait.png",
is_primary=True,
)
self.provider_mocks = {}
patch("apps.ai.services.get_image_provider", side_effect=self._provider_for).start()
patch("apps.ai.tasks.generate_triview_task.delay").start()
patch("apps.ai.services._store_generated_media", side_effect=self._store_media).start()
patch("apps.assets.review.submit_asset_for_review").start()
self.addCleanup(patch.stopall)
def provider(self, name, priority):
return ModelProvider.objects.create(
name=name,
display_name=name,
status=ModelProvider.Status.ACTIVE,
metadata={"routing": {"fallback_priority": priority}},
)
def model(self, provider, name, *, outbound=True, base_cost="0.50"):
return ModelConfig.objects.create(
provider=provider,
name=name,
display_name=name,
capability=ModelConfig.Capability.IMAGE,
endpoint="images/generations",
unit_price=Decimal("20"),
status=ModelConfig.Status.ACTIVE,
metadata=_metadata(outbound=outbound, base_cost=base_cost),
)
@staticmethod
def _new_provider_mock():
provider = Mock()
provider.image_edit.return_value = {"data": [{"url": "http://example.test/tri.png"}]}
provider.image_generation.return_value = {"data": [{"url": "http://example.test/tri.png"}]}
provider.extract_first_media_url.side_effect = lambda response: response["data"][0]["url"]
return provider
@staticmethod
def _new_generation_only_provider_mock():
provider = Mock(spec=["image_generation", "extract_first_media_url"])
provider.image_generation.return_value = {"data": [{"url": "http://example.test/tri.png"}]}
provider.extract_first_media_url.side_effect = lambda response: response["data"][0]["url"]
return provider
def _provider_for(self, model):
return self.provider_mocks.setdefault(model.id, self._new_provider_mock())
def _store_media(self, **kwargs):
return Asset.objects.create(
team=kwargs["team"],
created_by=kwargs["user"],
name=kwargs["name"],
asset_type=kwargs["asset_type"],
source=Asset.Source.AI_GENERATED,
category=kwargs["category"],
origin_task=kwargs["task"],
)
def submit(self):
return generate_person_triview(
project=self.project,
user=self.user,
portrait_asset=self.portrait,
)
@staticmethod
def ledger_count(task, ledger_type):
return CreditLedger.objects.filter(task=task, ledger_type=ledger_type).count()
def test_first_success_records_single_reference_16_9_and_keeps_project_scope(self):
primary = self.model(self.provider("project-tri-primary", 20), "project-tri-primary")
task = self.submit()
run_triview_task(task_id=str(task.id))
task.refresh_from_db()
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
self.assertTrue(task.request_payload["model_routing_v1"])
attempt = task.model_attempts.get()
self.assertEqual(attempt.model_config_id, primary.id)
self.assertEqual(attempt.operation, "image_edit")
self.assertEqual(attempt.request_summary["reference_images"], 1)
self.assertEqual(attempt.request_summary["aspect_ratio"], "16:9")
self.assertTrue(attempt.request_summary["project_person_triview"])
call = self.provider_mocks[primary.id].image_edit.call_args
self.assertEqual(call.kwargs["images"], ["http://example.test/portrait.png"])
self.assertEqual(call.kwargs["size"], "1536x864")
group = self.project.base_asset_groups.get(metadata__triview_of=str(self.portrait.id))
self.assertEqual(group.adopted_asset.origin_task_id, task.id)
self.assertFalse(Model.objects.filter(portrait_asset=self.portrait).exists())
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
def test_retry_then_dynamic_candidate_success_charges_once(self):
primary = self.model(
self.provider("project-tri-fallback-primary", 100),
"project-tri-primary",
base_cost="0.25",
)
candidate = self.model(
self.provider("volcano", 10),
"project-tri-candidate",
outbound=False,
base_cost="0.75",
)
task = self.submit()
self.provider_mocks[primary.id] = self._new_provider_mock()
self.provider_mocks[primary.id].image_edit.side_effect = requests.ConnectionError("offline")
run_triview_task(task_id=str(task.id))
task.refresh_from_db()
attempts = list(task.model_attempts.all())
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
self.assertEqual([a.model_config_id for a in attempts], [primary.id, primary.id, candidate.id])
self.assertTrue(attempts[-1].is_fallback)
self.assertTrue(all(a.request_summary["aspect_ratio"] == "16:9" for a in attempts))
self.assertEqual(task.base_cost, Decimal("0.7500"))
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
def test_all_candidates_fail_releases_and_keeps_existing_group(self):
primary = self.model(self.provider("project-tri-all-primary", 100), "project-tri-primary")
candidate = self.model(
self.provider("volcano", 10), "project-tri-candidate", outbound=False
)
old = Asset.objects.create(
team=self.team,
created_by=self.user,
name="旧三视图",
asset_type=Asset.Type.IMAGE,
source=Asset.Source.AI_GENERATED,
category=Asset.Category.TRI_VIEW,
)
group = BaseAssetGroup.objects.create(
project=self.project,
kind=BaseAssetGroup.Kind.PERSON,
prompt="旧提示词",
adopted_asset=old,
metadata={"label": "·三视图", "triview_of": str(self.portrait.id)},
)
group.candidate_assets.add(old)
task = self.submit()
for model in (primary, candidate):
self.provider_mocks[model.id] = self._new_provider_mock()
self.provider_mocks[model.id].image_edit.side_effect = requests.ConnectionError("offline")
run_triview_task(task_id=str(task.id))
task.refresh_from_db()
group.refresh_from_db()
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(task.model_attempts.count(), 3)
self.assertEqual(group.adopted_asset_id, old.id)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
def test_direct_primary_uses_image_generation_with_reference(self):
primary = self.model(
self.provider("volcano", 10), "seedream-project-triview", outbound=False
)
self.provider_mocks[primary.id] = self._new_generation_only_provider_mock()
task = self.submit()
run_triview_task(task_id=str(task.id))
task.refresh_from_db()
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
call = self.provider_mocks[primary.id].image_generation.call_args
self.assertEqual(call.kwargs["image"], ["http://example.test/portrait.png"])
self.assertEqual(call.kwargs["size"], "1536x864")
self.assertEqual(task.model_attempts.get().public_model_name, primary.display_name)
@@ -0,0 +1,366 @@
from decimal import Decimal
from unittest.mock import Mock
import requests
from django.test import TestCase
from apps.accounts.models import Team, TeamMember, User
from apps.ai.model_routing import ModelRequirements
from apps.ai.models import AIModelAttempt, AITask, ModelConfig, ModelProvider
from apps.ai.routing_executor import AttemptMetadata, execute_model_call, sanitize_attempt_summary
from apps.billing.models import CreditAccount, CreditLedger, CreditReservation
from apps.billing.services.ledger import charge_reserved_credit, release_credit, reserve_credit
def text_metadata(*, outbound=True, candidate=True):
return {
"routing": {"fallback_on_failure": outbound, "fallback_candidate": candidate},
"capabilities": {
"operations": ["chat"],
"features": ["streaming", "structured_output"],
},
}
class RoutingExecutorTests(TestCase):
requirements = ModelRequirements(capability="text", operation="chat")
def setUp(self):
# 迁移种子不会参与本组动态候选,保证断言只覆盖用例声明的模型。
ModelConfig.objects.all().update(status=ModelConfig.Status.DISABLED)
self.user = User.objects.create_user(username="routing-user", password="x")
self.team = Team.objects.create(name="Routing Team", owner=self.user)
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
CreditAccount.objects.create(team=self.team, balance=Decimal("1000"))
def provider(self, name, priority):
return ModelProvider.objects.create(
name=name,
display_name=name,
metadata={"routing": {"fallback_priority": priority}},
)
def model(self, provider, name, *, outbound=True, candidate=True):
return ModelConfig.objects.create(
provider=provider,
name=name,
display_name=name,
capability=ModelConfig.Capability.TEXT,
endpoint="chat/completions",
metadata=text_metadata(outbound=outbound, candidate=candidate),
)
def task(self, model, key):
return AITask.objects.create(
team=self.team,
created_by=self.user,
model_config=model,
task_type=AITask.Type.SCRIPT_GENERATION,
status=AITask.Status.RESERVED,
idempotency_key=key,
estimated_cost=Decimal("10"),
actual_cost=Decimal("10"),
)
def execute(self, *, task, primary, invoke, **kwargs):
return execute_model_call(
task=task,
primary_model=primary,
requirements=self.requirements,
public_model_name="AirShelf Script",
invoke=invoke,
sleep=lambda seconds: None,
uniform=lambda start, end: 0,
**kwargs,
)
def test_first_success_logs_once_and_never_queries_candidates(self):
primary = self.model(self.provider("exec-first", 100), "primary")
task = self.task(primary, "exec-first")
resolver = Mock(return_value=[])
result = self.execute(
task=task,
primary=primary,
invoke=lambda model, timeout: {"ok": True},
candidate_resolver=resolver,
)
self.assertEqual(result.actual_model, primary)
self.assertFalse(result.fallback_used)
self.assertEqual(result.call_count, 1)
resolver.assert_not_called()
attempt = task.model_attempts.get()
self.assertEqual(attempt.status, AIModelAttempt.Status.SUCCEEDED)
self.assertFalse(attempt.is_retry)
self.assertFalse(attempt.is_fallback)
def test_transient_failure_retries_primary_then_succeeds(self):
primary = self.model(self.provider("exec-retry", 100), "primary")
task = self.task(primary, "exec-retry")
resolver = Mock(return_value=[])
calls = []
def invoke(model, timeout):
calls.append((model.id, timeout))
if len(calls) == 1:
raise requests.ConnectionError("temporary network failure")
return {"ok": True}
result = self.execute(
task=task,
primary=primary,
invoke=invoke,
candidate_resolver=resolver,
)
self.assertEqual(result.call_count, 2)
resolver.assert_not_called()
attempts = list(task.model_attempts.all())
self.assertEqual([item.status for item in attempts], ["failed", "succeeded"])
self.assertEqual([item.is_retry for item in attempts], [False, True])
self.assertEqual(attempts[1].previous_attempt_id, attempts[0].id)
def test_non_retryable_provider_failure_switches_and_preserves_primary(self):
primary = self.model(self.provider("exec-fallback-primary", 100), "primary")
candidate = self.model(self.provider("exec-fallback-candidate", 10), "candidate")
task = self.task(primary, "exec-fallback")
resolver = Mock(return_value=[candidate])
def invoke(model, timeout):
if model.id == primary.id:
raise RuntimeError("API_KEY is not configured")
return {"ok": True}
result = self.execute(
task=task,
primary=primary,
invoke=invoke,
candidate_resolver=resolver,
)
self.assertEqual(result.actual_model, candidate)
self.assertTrue(result.fallback_used)
task.refresh_from_db()
self.assertEqual(task.model_config_id, primary.id)
attempts = list(task.model_attempts.all())
self.assertEqual([item.model_name for item in attempts], ["primary", "candidate"])
self.assertEqual([item.public_model_name for item in attempts], ["AirShelf Script", "AirShelf Script"])
self.assertTrue(attempts[1].is_fallback)
excluded = resolver.call_args.kwargs["excluded_provider_ids"]
self.assertIn(primary.provider_id, excluded)
def test_all_models_fail_without_cycles_and_stops_at_global_call_limit(self):
primary = self.model(self.provider("exec-all-primary", 500), "primary")
self.model(self.provider("exec-all-first", 10), "candidate-1")
self.model(self.provider("exec-all-second", 20), "candidate-2")
task = self.task(primary, "exec-all")
with self.assertRaises(requests.ConnectionError):
self.execute(
task=task,
primary=primary,
invoke=lambda model, timeout: (_ for _ in ()).throw(requests.ConnectionError("offline")),
)
attempts = list(task.model_attempts.all())
self.assertEqual(len(attempts), 5)
self.assertEqual([item.model_name for item in attempts], ["primary", "primary", "primary", "candidate-1", "candidate-2"])
self.assertEqual(len({item.sequence for item in attempts}), 5)
self.assertTrue(all(item.status == AIModelAttempt.Status.FAILED for item in attempts))
def test_content_rejection_neither_retries_nor_fallbacks(self):
primary = self.model(self.provider("exec-safety", 100), "primary")
task = self.task(primary, "exec-safety")
resolver = Mock(return_value=[])
with self.assertRaises(RuntimeError):
self.execute(
task=task,
primary=primary,
invoke=lambda model, timeout: (_ for _ in ()).throw(RuntimeError("moderation_blocked")),
candidate_resolver=resolver,
)
resolver.assert_not_called()
attempt = task.model_attempts.get()
self.assertEqual(attempt.error_type, "content_rejected")
def test_retry_after_above_cap_skips_wait_and_switches_immediately(self):
primary = self.model(self.provider("exec-rate-primary", 100), "primary")
candidate = self.model(self.provider("exec-rate-candidate", 10), "candidate")
task = self.task(primary, "exec-rate")
sleeps = []
def invoke(model, timeout):
if model.id == primary.id:
response = requests.Response()
response.status_code = 429
response.headers["Retry-After"] = "60"
exc = requests.HTTPError("rate limited", response=response)
raise exc
return {"ok": True}
result = execute_model_call(
task=task,
primary_model=primary,
requirements=self.requirements,
public_model_name="AirShelf Script",
invoke=invoke,
candidate_resolver=Mock(return_value=[candidate]),
sleep=sleeps.append,
uniform=lambda start, end: 0,
)
self.assertEqual(result.call_count, 2)
self.assertEqual(sleeps, [])
def test_video_provider_task_id_on_error_never_retries_or_fallbacks(self):
provider = self.provider("exec-video", 100)
model = ModelConfig.objects.create(
provider=provider,
name="video-model",
display_name="video-model",
capability=ModelConfig.Capability.VIDEO,
endpoint="contents/generations/tasks",
metadata={
"routing": {"fallback_on_failure": True, "fallback_candidate": True},
"capabilities": {
"operations": ["video_generate"],
"features": ["text_to_video"],
"resolutions": ["720p"],
"durations": [5],
},
"pricing": {"default": {"no_ref_video": 1}},
},
)
task = self.task(model, "exec-video-known-id")
resolver = Mock(return_value=[])
with self.assertRaises(requests.Timeout):
execute_model_call(
task=task,
primary_model=model,
requirements=ModelRequirements(
capability="video",
operation="video_generate",
resolution="720p",
duration=5,
),
public_model_name="Seedance",
invoke=lambda current, timeout: (_ for _ in ()).throw(requests.Timeout("poll state unknown")),
error_metadata=lambda exc, current: AttemptMetadata(provider_task_id="provider-task-123"),
candidate_resolver=resolver,
sleep=lambda seconds: None,
)
resolver.assert_not_called()
self.assertEqual(task.model_attempts.count(), 1)
self.assertEqual(task.model_attempts.get().provider_task_id, "provider-task-123")
def test_metadata_extraction_failure_does_not_repeat_successful_model_call(self):
primary = self.model(self.provider("exec-meta", 100), "primary")
task = self.task(primary, "exec-meta")
invoke = Mock(return_value={"ok": True})
result = self.execute(
task=task,
primary=primary,
invoke=invoke,
result_metadata=lambda value, model: (_ for _ in ()).throw(RuntimeError("summary failed")),
)
self.assertEqual(result.call_count, 1)
invoke.assert_called_once()
self.assertIn("metadata_error", task.model_attempts.get().response_summary)
def test_attempt_costs_are_audited_and_summed_without_touching_parent_updated_at(self):
primary = self.model(self.provider("exec-cost-primary", 100), "primary")
candidate = self.model(self.provider("exec-cost-candidate", 10), "candidate")
task = self.task(primary, "exec-cost")
original_updated_at = task.updated_at
def invoke(model, timeout):
if model.id == primary.id:
exc = RuntimeError("API_KEY is not configured")
exc.attempt_metadata = {"platform_cost": "0.125", "usage": {"prompt_tokens": 10}}
raise exc
return {"ok": True}
self.execute(
task=task,
primary=primary,
invoke=invoke,
candidate_resolver=Mock(return_value=[candidate]),
result_metadata=lambda result, model: AttemptMetadata(
platform_cost=Decimal("0.375"), usage={"output_tokens": 20}
),
)
task.refresh_from_db()
self.assertEqual(task.base_cost, Decimal("0.5000"))
self.assertEqual(task.updated_at, original_updated_at)
self.assertEqual(
list(task.model_attempts.values_list("platform_cost", flat=True)),
[Decimal("0.1250"), Decimal("0.3750")],
)
def test_feedback_summary_redacts_credentials_and_large_binary(self):
summary = sanitize_attempt_summary(
{
"api_key": "secret-value",
"authorization": "Bearer secret",
"prompt": "hello",
"image": b"123456",
}
)
self.assertEqual(summary["api_key"], "[REDACTED]")
self.assertEqual(summary["authorization"], "[REDACTED]")
self.assertEqual(summary["image"], "[二进制 6 字节]")
def test_fallback_success_keeps_single_reserve_and_single_charge(self):
primary = self.model(self.provider("exec-bill-primary", 100), "primary")
candidate = self.model(self.provider("exec-bill-candidate", 10), "candidate")
task = self.task(primary, "exec-bill-success")
reservation = reserve_credit(team=self.team, user=self.user, task=task, amount=Decimal("10"))
def invoke(model, timeout):
if model.id == primary.id:
raise RuntimeError("API_KEY is not configured")
return {"ok": True}
self.execute(
task=task,
primary=primary,
invoke=invoke,
candidate_resolver=Mock(return_value=[candidate]),
)
charge_reserved_credit(reservation=reservation, actual_amount=Decimal("10"))
self.assertEqual(CreditReservation.objects.filter(task=task).count(), 1)
self.assertEqual(CreditLedger.objects.filter(task=task, ledger_type=CreditLedger.Type.RESERVE).count(), 1)
self.assertEqual(CreditLedger.objects.filter(task=task, ledger_type=CreditLedger.Type.CHARGE).count(), 1)
self.assertEqual(task.model_attempts.count(), 2)
def test_all_failures_keep_single_reserve_and_single_release(self):
primary = self.model(self.provider("exec-release-primary", 500), "primary")
self.model(self.provider("exec-release-first", 10), "candidate-1")
self.model(self.provider("exec-release-second", 20), "candidate-2")
task = self.task(primary, "exec-bill-failure")
reservation = reserve_credit(team=self.team, user=self.user, task=task, amount=Decimal("10"))
with self.assertRaises(requests.ConnectionError):
self.execute(
task=task,
primary=primary,
invoke=lambda model, timeout: (_ for _ in ()).throw(requests.ConnectionError("offline")),
)
release_credit(reservation=reservation, reason="all providers failed")
self.assertEqual(CreditReservation.objects.filter(task=task).count(), 1)
self.assertEqual(CreditLedger.objects.filter(task=task, ledger_type=CreditLedger.Type.RESERVE).count(), 1)
self.assertEqual(CreditLedger.objects.filter(task=task, ledger_type=CreditLedger.Type.CHARGE).count(), 0)
self.assertEqual(CreditLedger.objects.filter(task=task, ledger_type=CreditLedger.Type.RELEASE).count(), 1)
account = CreditAccount.objects.get(team=self.team)
self.assertEqual(account.reserved_balance, Decimal("0"))
+157
View File
@@ -0,0 +1,157 @@
from copy import deepcopy
from dataclasses import FrozenInstanceError
import os
from unittest.mock import patch
from django.apps import apps
from django.conf import settings
from django.core.exceptions import ImproperlyConfigured
from django.test import SimpleTestCase, override_settings
from apps.ai.routing_policy import RoutingPolicyConfigurationError, load_model_routing_policy
from airshelf.settings.base import env_float, env_int, env_int_list
class ModelRoutingPolicyTests(SimpleTestCase):
def policy_dict(self) -> dict:
return deepcopy(settings.MODEL_ROUTING_POLICY)
def test_default_policy_matches_confirmed_values_and_is_immutable(self):
policy = load_model_routing_policy()
self.assertEqual(policy.max_models, 3)
self.assertEqual(policy.max_calls, 5)
self.assertEqual(policy.jitter_ratio, 0.20)
self.assertEqual(policy.text.retry_delays, (1.0, 3.0))
self.assertEqual(policy.text.request_timeout, 120.0)
self.assertEqual(policy.text.stream_timeout, 300.0)
self.assertEqual(policy.text.total_timeout, 480.0)
self.assertEqual(policy.image.retry_delays, (3.0,))
self.assertEqual(policy.image.request_timeout, 300.0)
self.assertEqual(policy.image.total_timeout, 900.0)
self.assertEqual(policy.audio.retry_delays, (2.0,))
self.assertEqual(policy.audio.total_timeout, 180.0)
self.assertEqual(policy.video.submit_retry_delays, (3.0,))
self.assertEqual(policy.video.submit_timeout, 120.0)
self.assertEqual(policy.video.submit_total_timeout, 300.0)
self.assertEqual(policy.video.poll_request_timeout, 60.0)
self.assertEqual(policy.video.generation_timeout, 1800.0)
self.assertEqual(policy.postprocess.retry_delays, (2.0, 5.0, 10.0))
with self.assertRaises(FrozenInstanceError):
policy.max_calls = 6
def test_settings_can_inject_short_timeouts_without_production_branch(self):
raw = self.policy_dict()
raw["text"].update(
retry_delays=[0],
request_timeout=1,
stream_timeout=1,
total_timeout=2,
retry_after_cap=0,
)
raw["image"].update(retry_delays=[], request_timeout=1, total_timeout=1, retry_after_cap=0)
raw["audio"].update(retry_delays=[], request_timeout=1, total_timeout=1, retry_after_cap=0)
raw["video"].update(
submit_retry_delays=[],
submit_timeout=1,
submit_total_timeout=1,
poll_request_timeout=1,
generation_timeout=1,
retry_after_cap=0,
)
raw["postprocess"]["retry_delays"] = []
with override_settings(MODEL_ROUTING_POLICY=raw):
policy = load_model_routing_policy()
self.assertEqual(policy.text.retry_delays, (0.0,))
self.assertEqual(policy.text.total_timeout, 2.0)
self.assertEqual(policy.image.retry_delays, ())
self.assertEqual(policy.video.generation_timeout, 1.0)
self.assertEqual(policy.postprocess.retry_delays, ())
def test_rejects_max_calls_smaller_than_model_count_with_chinese_message(self):
raw = self.policy_dict()
raw["max_models"] = 3
raw["max_calls"] = 2
with self.assertRaisesRegex(
RoutingPolicyConfigurationError,
"最多真实模型调用次数.*不得小于 max_models.*最多尝试模型数量",
):
load_model_routing_policy(raw)
def test_rejects_single_timeout_greater_than_total_timeout(self):
raw = self.policy_dict()
raw["image"]["request_timeout"] = 901
raw["image"]["total_timeout"] = 900
with self.assertRaisesRegex(
RoutingPolicyConfigurationError,
"图片|request_timeout.*单次调用超时.*不得大于.*total_timeout",
):
load_model_routing_policy(raw)
def test_rejects_non_numeric_retry_delay_and_reports_value(self):
raw = self.policy_dict()
raw["audio"]["retry_delays"] = ["两秒"]
with self.assertRaisesRegex(
RoutingPolicyConfigurationError,
r"audio\.retry_delays\[0\].*必须是数字.*两秒",
):
load_model_routing_policy(raw)
def test_rejects_missing_required_section(self):
raw = self.policy_dict()
del raw["postprocess"]
with self.assertRaisesRegex(
RoutingPolicyConfigurationError,
"MODEL_ROUTING_POLICY.postprocess.*必填配置",
):
load_model_routing_policy(raw)
def test_rejects_video_poll_timeout_longer_than_generation_window(self):
raw = self.policy_dict()
raw["video"]["poll_request_timeout"] = 61
raw["video"]["generation_timeout"] = 60
with self.assertRaisesRegex(
RoutingPolicyConfigurationError,
"视频单次轮询超时.*不得大于.*generation_timeout",
):
load_model_routing_policy(raw)
def test_app_startup_fails_fast_for_invalid_policy(self):
raw = self.policy_dict()
raw["jitter_ratio"] = 1.5
with override_settings(MODEL_ROUTING_POLICY=raw):
with self.assertRaisesRegex(ImproperlyConfigured, "模型路由策略配置错误.*重试随机抖动比例"):
apps.get_app_config("ai").ready()
class RoutingPolicyEnvironmentParserTests(SimpleTestCase):
def test_integer_float_and_list_environment_overrides(self):
with patch.dict(
os.environ,
{
"TEST_ROUTING_INT": "7",
"TEST_ROUTING_FLOAT": "0.25",
"TEST_ROUTING_LIST": "1, 3,5",
},
):
self.assertEqual(env_int("TEST_ROUTING_INT", 1), 7)
self.assertEqual(env_float("TEST_ROUTING_FLOAT", 0.1), 0.25)
self.assertEqual(env_int_list("TEST_ROUTING_LIST", (9,)), [1, 3, 5])
def test_empty_retry_delay_environment_means_no_retry(self):
with patch.dict(os.environ, {"TEST_ROUTING_LIST": ""}):
self.assertEqual(env_int_list("TEST_ROUTING_LIST", (1, 3)), [])
def test_invalid_environment_value_has_chinese_error(self):
with patch.dict(os.environ, {"TEST_ROUTING_INT": ""}):
with self.assertRaisesRegex(ValueError, "环境变量 TEST_ROUTING_INT 必须是整数.*三"):
env_int("TEST_ROUTING_INT", 1)
@@ -0,0 +1,179 @@
from decimal import Decimal
from unittest.mock import Mock, patch
from django.test import TestCase
from apps.accounts.models import Team, TeamMember, User
from apps.ai.models import AITask, ModelConfig, ModelProvider
from apps.ai.script_agent import persist_script_draft
from apps.ai.services import extract_cast_and_scenes
from apps.billing.models import CreditAccount, CreditLedger
from apps.products.models import Product
from apps.projects.models import Project
class ScriptEntitySyncTests(TestCase):
def setUp(self):
ModelConfig.objects.filter(capability=ModelConfig.Capability.TEXT).update(
status=ModelConfig.Status.DISABLED
)
self.user = User.objects.create_user(username="script-entity-sync", password="x")
self.team = Team.objects.create(name="Script Entity Sync", owner=self.user)
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
CreditAccount.objects.create(team=self.team, balance=Decimal("1000"))
self.product = Product.objects.create(
team=self.team,
created_by=self.user,
title="测试商品",
)
self.project = Project.objects.create(
team=self.team,
created_by=self.user,
product=self.product,
name="脚本实体同步项目",
)
provider = ModelProvider.objects.create(
name="script-sync-provider",
display_name="Script Sync Provider",
status=ModelProvider.Status.ACTIVE,
)
self.model = ModelConfig.objects.create(
provider=provider,
name="script-sync-model",
display_name="Script Sync Model",
capability=ModelConfig.Capability.TEXT,
endpoint="chat/completions",
unit_price=Decimal("10"),
status=ModelConfig.Status.ACTIVE,
metadata={
"routing": {"fallback_on_failure": True, "fallback_candidate": True},
"capabilities": {
"operations": ["chat"],
"features": ["structured_output"],
},
"pricing": {"base_cost_yuan": "0.50"},
},
)
def main_task(self):
return AITask.objects.create(
team=self.team,
created_by=self.user,
project=self.project,
task_type=AITask.Type.SCRIPT_GENERATION,
status=AITask.Status.SUBMITTED,
model_config=self.model,
idempotency_key=f"script-entity-sync:{self.project.id}",
)
@staticmethod
def draft():
return {
"hook": "开场钩子",
"tone": "自然",
"aspect_ratio": "9:16",
"total_duration": 15,
"segment_count": 1,
"entities": [
{
"id": "c1",
"type": "character",
"name": "女主",
"visual_prompt": "都市女主",
"ref_index": 1,
},
{
"id": "s1",
"type": "scene",
"name": "客厅",
"visual_prompt": "现代客厅",
"ref_index": 2,
},
],
"segments": [
{
"index": 0,
"duration": 15,
"narration": "女主在客厅展示商品",
"visual": "女主站在客厅",
"role": "钩子",
"speaker": "女主",
"product_exposure": "展示",
"entity_refs": ["c1", "s1"],
"dialogue": [],
}
],
}
def test_script_draft_syncs_entities_without_creating_second_ai_task(self):
task = self.main_task()
before_count = AITask.objects.count()
script = persist_script_draft(
project=self.project,
user=self.user,
task=task,
draft=self.draft(),
source="ai",
)
self.project.refresh_from_db()
task.refresh_from_db()
self.assertEqual(AITask.objects.count(), before_count)
self.assertEqual(task.status, AITask.Status.SUBMITTED)
self.assertEqual(script.task_id, task.id)
self.assertEqual(self.project.metadata["cast"], ["女主"])
self.assertEqual(self.project.metadata["scenes"], ["客厅"])
self.assertEqual(
[entity["id"] for entity in self.project.metadata["script_entities"]],
["c1", "s1"],
)
self.assertEqual(CreditLedger.objects.filter(task=task).count(), 0)
@patch("apps.ai.services.get_text_provider")
def test_legacy_post_script_extractor_uses_routed_attempt_and_one_settlement(self, get_provider):
provider = Mock()
response = {
"choices": [
{
"message": {
"content": (
'{"cast":[{"name":"女主","prompt":"都市女主"}],'
'"scenes":[{"name":"客厅","prompt":"现代客厅"}]}'
)
}
}
],
"usage": {"total_tokens": 42},
}
provider.chat_completion.return_value = response
provider.extract_text.return_value = response["choices"][0]["message"]["content"]
get_provider.return_value = provider
result = extract_cast_and_scenes(
project=self.project,
user=self.user,
content="女主在客厅展示商品",
)
task = AITask.objects.get(task_type=AITask.Type.SCRIPT_OPTIMIZATION)
self.assertEqual(result["cast"], ["女主"])
self.assertEqual(result["scenes"], ["客厅"])
self.assertTrue(task.request_payload["model_routing_v1"])
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
attempt = task.model_attempts.get()
self.assertEqual(attempt.operation, "chat")
self.assertEqual(attempt.public_model_name, "AirShelf Script")
self.assertEqual(attempt.request_summary["source"], "legacy_post_script_extract")
self.assertEqual(
CreditLedger.objects.filter(task=task, ledger_type=CreditLedger.Type.RESERVE).count(),
1,
)
self.assertEqual(
CreditLedger.objects.filter(task=task, ledger_type=CreditLedger.Type.CHARGE).count(),
1,
)
self.assertEqual(
CreditLedger.objects.filter(task=task, ledger_type=CreditLedger.Type.RELEASE).count(),
0,
)
@@ -0,0 +1,226 @@
import json
from decimal import Decimal
from unittest.mock import Mock, patch
import requests
from django.test import TestCase
from apps.accounts.models import Team, TeamMember, User
from apps.ai.models import AITask, ModelConfig, ModelProvider
from apps.ai.script_agent import regenerate_segment_via_agent
from apps.billing.models import CreditAccount, CreditLedger
from apps.products.models import Product
from apps.projects.models import Project, ScriptSegment, ScriptVersion
def _metadata(*, outbound=True, base_cost="0.50"):
return {
"routing": {"fallback_on_failure": outbound, "fallback_candidate": True},
"capabilities": {
"operations": ["chat"],
"features": ["structured_output"],
},
"pricing": {"base_cost_yuan": base_cost},
}
class ScriptOptimizationRoutingTests(TestCase):
def setUp(self):
ModelConfig.objects.filter(capability=ModelConfig.Capability.TEXT).update(
status=ModelConfig.Status.DISABLED
)
self.user = User.objects.create_user(username="script-opt-routing", password="x")
self.team = Team.objects.create(name="Script Opt Routing", owner=self.user)
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
CreditAccount.objects.create(team=self.team, balance=Decimal("1000"))
product = Product.objects.create(team=self.team, created_by=self.user, title="测试商品")
self.project = Project.objects.create(
team=self.team,
created_by=self.user,
product=product,
name="单镜优化路由项目",
)
self.base_draft = {
"hook": "钩子",
"tone": "自然",
"aspect_ratio": "9:16",
"total_duration": 30,
"segment_count": 2,
"entities": [],
"segments": [
{
"index": 0,
"duration": 15,
"role": "钩子",
"narration": "旧口播0",
"visual": "旧画面0",
"speaker": None,
"product_exposure": "",
"entity_refs": [],
"dialogue": [],
},
{
"index": 1,
"duration": 15,
"role": "CTA",
"narration": "旧口播1",
"visual": "旧画面1",
"speaker": None,
"product_exposure": "",
"entity_refs": [],
"dialogue": [],
},
],
}
self.script = ScriptVersion.objects.create(
project=self.project,
title="基准稿",
content=json.dumps(self.base_draft, ensure_ascii=False),
is_adopted=True,
metadata={
"hook": "钩子",
"tone": "自然",
"aspect_ratio": "9:16",
"total_duration": 30,
"entities": [],
},
)
ScriptSegment.objects.create(
script_version=self.script,
sort_order=0,
duration_seconds=15,
role="钩子",
narration="旧口播0",
visual_prompt="旧画面0",
)
self.target = ScriptSegment.objects.create(
script_version=self.script,
sort_order=1,
duration_seconds=15,
role="CTA",
narration="旧口播1",
visual_prompt="旧画面1",
)
self.provider_mocks = {}
patch("apps.ai.services.get_text_provider", side_effect=self._provider_for).start()
self.addCleanup(patch.stopall)
def provider(self, name, priority):
return ModelProvider.objects.create(
name=name,
display_name=name,
status=ModelProvider.Status.ACTIVE,
metadata={"routing": {"fallback_priority": priority}},
)
def model(self, provider, name, *, outbound=True, base_cost="0.50"):
return ModelConfig.objects.create(
provider=provider,
name=name,
display_name=name,
capability=ModelConfig.Capability.TEXT,
endpoint="chat/completions",
unit_price=Decimal("10"),
status=ModelConfig.Status.ACTIVE,
metadata=_metadata(outbound=outbound, base_cost=base_cost),
)
def valid_text(self):
draft = json.loads(json.dumps(self.base_draft))
draft["segments"][1]["narration"] = "优化后的口播"
draft["segments"][1]["visual"] = "优化后的画面"
return json.dumps(draft, ensure_ascii=False)
def good_provider(self):
provider = Mock()
provider.chat_completion.return_value = {"usage": {"total_tokens": 42}}
provider.extract_text.return_value = self.valid_text()
return provider
def _provider_for(self, model):
return self.provider_mocks.setdefault(model.id, self.good_provider())
@staticmethod
def ledger_count(task, ledger_type):
return CreditLedger.objects.filter(task=task, ledger_type=ledger_type).count()
def optimize(self, primary):
return regenerate_segment_via_agent(
project=self.project,
user=self.user,
model_config=primary,
segment=self.target,
instruction="让第二镜更有吸引力",
)
def test_first_success_records_attempt_and_preserves_other_segment(self):
primary = self.model(self.provider("script-opt-primary", 20), "script-opt-primary")
version = self.optimize(primary)
task = AITask.objects.get(task_type=AITask.Type.SCRIPT_OPTIMIZATION)
segments = list(version.segments.order_by("sort_order"))
self.assertEqual(segments[0].narration, "旧口播0")
self.assertEqual(segments[1].narration, "优化后的口播")
self.assertTrue(task.request_payload["model_routing_v1"])
attempt = task.model_attempts.get()
self.assertEqual(attempt.model_config_id, primary.id)
self.assertEqual(attempt.operation, "chat")
self.assertEqual(attempt.request_summary["target_index"], 1)
self.assertEqual(attempt.public_model_name, "AirShelf Script")
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
def test_primary_retries_then_dynamic_fallback_succeeds_once(self):
primary = self.model(
self.provider("script-opt-fallback-primary", 100),
"script-opt-primary",
base_cost="0.25",
)
candidate = self.model(
self.provider("script-opt-fallback-candidate", 10),
"script-opt-candidate",
outbound=False,
base_cost="0.75",
)
failed = self.good_provider()
failed.chat_completion.side_effect = requests.ConnectionError("offline")
self.provider_mocks[primary.id] = failed
version = self.optimize(primary)
task = AITask.objects.get(task_type=AITask.Type.SCRIPT_OPTIMIZATION)
attempts = list(task.model_attempts.all())
self.assertEqual(
[attempt.model_config_id for attempt in attempts],
[primary.id, primary.id, primary.id, candidate.id],
)
self.assertTrue(attempts[-1].is_fallback)
self.assertEqual(version.segments.order_by("sort_order")[0].narration, "旧口播0")
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
def test_all_fail_releases_once_and_does_not_create_new_version(self):
primary = self.model(self.provider("script-opt-all-primary", 100), "script-opt-primary")
fallback = self.model(
self.provider("script-opt-all-fallback", 10),
"script-opt-fallback",
outbound=False,
)
for model in (primary, fallback):
failed = self.good_provider()
failed.chat_completion.side_effect = requests.ConnectionError("offline")
self.provider_mocks[model.id] = failed
before_versions = ScriptVersion.objects.filter(project=self.project).count()
with self.assertRaises(Exception):
self.optimize(primary)
task = AITask.objects.get(task_type=AITask.Type.SCRIPT_OPTIMIZATION)
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(ScriptVersion.objects.filter(project=self.project).count(), before_versions)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
@@ -0,0 +1,349 @@
import json
import time
from decimal import Decimal
from unittest.mock import Mock, patch
import requests
from django.test import TransactionTestCase
from apps.accounts.models import Team, TeamMember, User
from apps.ai.models import AITask, ModelConfig, ModelProvider
from apps.ai.script_agent import stream_script_agent
from apps.billing.models import CreditAccount, CreditLedger
from apps.products.models import Product
from apps.projects.models import Project, ScriptSegment, ScriptVersion
def _metadata(*, outbound=True, base_cost="0.50"):
return {
"routing": {"fallback_on_failure": outbound, "fallback_candidate": True},
"capabilities": {
"operations": ["chat"],
"features": ["streaming", "structured_output"],
},
"pricing": {"base_cost_yuan": base_cost},
}
class ScriptStreamRoutingTests(TransactionTestCase):
reset_sequences = True
def setUp(self):
ModelConfig.objects.filter(capability=ModelConfig.Capability.TEXT).update(
status=ModelConfig.Status.DISABLED
)
self.user = User.objects.create_user(username="script-stream-routing", password="x")
self.team = Team.objects.create(name="Script Stream Routing", owner=self.user)
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
CreditAccount.objects.create(team=self.team, balance=Decimal("1000"))
product = Product.objects.create(team=self.team, created_by=self.user, title="测试商品")
self.project = Project.objects.create(
team=self.team,
created_by=self.user,
product=product,
name="流式脚本路由项目",
)
self.provider_mocks = {}
patch("apps.ai.services.get_text_provider", side_effect=self._provider_for).start()
self.addCleanup(patch.stopall)
def provider(self, name, priority):
return ModelProvider.objects.create(
name=name,
display_name=name,
status=ModelProvider.Status.ACTIVE,
metadata={"routing": {"fallback_priority": priority}},
)
def model(self, provider, name, *, outbound=True, base_cost="0.50"):
return ModelConfig.objects.create(
provider=provider,
name=name,
display_name=name,
capability=ModelConfig.Capability.TEXT,
endpoint="chat/completions",
unit_price=Decimal("10"),
status=ModelConfig.Status.ACTIVE,
metadata=_metadata(outbound=outbound, base_cost=base_cost),
)
@staticmethod
def valid_draft():
return {
"hook": "开场钩子",
"tone": "自然",
"aspect_ratio": "9:16",
"total_duration": 15,
"segment_count": 1,
"entities": [
{
"id": "c1",
"type": "character",
"name": "女主",
"visual_prompt": "都市女主",
"ref_index": 1,
}
],
"segments": [
{
"index": 0,
"duration": 15,
"role": "钩子",
"narration": "全新脚本口播",
"visual": "女主展示商品",
"speaker": "女主",
"product_exposure": "展示",
"entity_refs": ["c1"],
"dialogue": [],
}
],
}
def good_provider(self):
provider = Mock()
raw = json.dumps(self.valid_draft(), ensure_ascii=False)
provider.chat_completion_stream.return_value = iter(
[
{"type": "reasoning", "text": "先分析商品卖点"},
{"type": "delta", "text": "我先给你一版。\n"},
{"type": "delta", "text": raw},
{"type": "done"},
]
)
return provider
def _provider_for(self, model):
return self.provider_mocks.setdefault(model.id, self.good_provider())
@staticmethod
def parse_events(frames):
return [json.loads(frame.removeprefix("data: ").strip()) for frame in frames]
@staticmethod
def ledger_count(task, ledger_type):
return CreditLedger.objects.filter(task=task, ledger_type=ledger_type).count()
def generate(self, primary):
return self.parse_events(
list(
stream_script_agent(
project=self.project,
user=self.user,
model_config=primary,
mode="auto",
user_prompt="生成一版脚本",
aspect_ratio="9:16",
total_duration=15,
)
)
)
def test_first_success_keeps_reasoning_and_delta_sse_and_records_attempt(self):
primary = self.model(self.provider("script-stream-primary", 20), "script-stream-primary")
events = self.generate(primary)
task = AITask.objects.get(task_type=AITask.Type.SCRIPT_GENERATION)
self.assertIn("reasoning", [event["type"] for event in events])
self.assertIn("delta", [event["type"] for event in events])
self.assertIn("draft", [event["type"] for event in events])
self.assertEqual(events[-1]["type"], "done")
self.assertEqual(ScriptVersion.objects.filter(project=self.project).count(), 1)
attempt = task.model_attempts.get()
self.assertEqual(attempt.model_config_id, primary.id)
self.assertTrue(attempt.request_summary["streaming"])
self.assertTrue(attempt.request_summary["structured_output"])
self.assertEqual(attempt.request_summary["business_operation"], "script_generate")
self.assertEqual(attempt.public_model_name, "AirShelf Script")
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
def test_primary_retries_then_fallback_succeeds_with_one_script_and_settlement(self):
primary = self.model(
self.provider("script-stream-fallback-primary", 100),
"script-stream-primary",
base_cost="0.25",
)
candidate = self.model(
self.provider("script-stream-fallback-candidate", 10),
"script-stream-candidate",
outbound=False,
base_cost="0.75",
)
failed = Mock()
failed.chat_completion_stream.side_effect = requests.ConnectionError("offline")
self.provider_mocks[primary.id] = failed
events = self.generate(primary)
task = AITask.objects.get(task_type=AITask.Type.SCRIPT_GENERATION)
attempts = list(task.model_attempts.all())
self.assertEqual(events[-1]["type"], "done")
self.assertEqual(ScriptVersion.objects.filter(project=self.project).count(), 1)
self.assertEqual(
[attempt.model_config_id for attempt in attempts],
[primary.id, primary.id, primary.id, candidate.id],
)
self.assertTrue(attempts[-1].is_fallback)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
def test_all_fail_emits_safe_error_releases_once_and_saves_no_script(self):
primary = self.model(self.provider("script-stream-all-primary", 100), "script-stream-primary")
fallback = self.model(
self.provider("script-stream-all-fallback", 10),
"script-stream-fallback",
outbound=False,
)
for model in (primary, fallback):
failed = Mock()
failed.chat_completion_stream.side_effect = requests.ConnectionError("offline")
self.provider_mocks[model.id] = failed
events = self.generate(primary)
task = AITask.objects.get(task_type=AITask.Type.SCRIPT_GENERATION)
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(events[-1]["type"], "error")
self.assertNotIn("offline", events[-1]["detail"])
self.assertEqual(ScriptVersion.objects.filter(project=self.project).count(), 0)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
def test_client_disconnect_stops_stream_and_releases_without_saving_script(self):
primary = self.model(self.provider("script-stream-abort", 20), "script-stream-primary")
raw = json.dumps(self.valid_draft(), ensure_ascii=False)
def delayed_events():
yield {"type": "reasoning", "text": "正在分析"}
time.sleep(0.05)
yield {"type": "delta", "text": raw}
yield {"type": "done"}
provider = Mock()
provider.chat_completion_stream.return_value = delayed_events()
self.provider_mocks[primary.id] = provider
stream = stream_script_agent(
project=self.project,
user=self.user,
model_config=primary,
mode="auto",
user_prompt="生成一版脚本",
aspect_ratio="9:16",
total_duration=15,
)
for frame in stream:
event = json.loads(frame.removeprefix("data: ").strip())
if event["type"] == "reasoning":
stream.close()
break
deadline = time.monotonic() + 1
task = AITask.objects.get(task_type=AITask.Type.SCRIPT_GENERATION)
while not task.model_attempts.filter(status="failed").exists() and time.monotonic() < deadline:
time.sleep(0.01)
task.refresh_from_db()
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(ScriptVersion.objects.filter(project=self.project).count(), 0)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
def test_stream_revise_target_keeps_other_segment_and_uses_base_duration(self):
primary = self.model(self.provider("script-stream-revise", 20), "script-stream-primary")
base_draft = {
"hook": "原钩子",
"tone": "自然",
"aspect_ratio": "9:16",
"total_duration": 30,
"segment_count": 2,
"entities": [],
"segments": [
{
"index": 0,
"duration": 15,
"role": "钩子",
"narration": "旧口播0",
"visual": "旧画面0",
"speaker": None,
"product_exposure": "",
"entity_refs": [],
"dialogue": [],
},
{
"index": 1,
"duration": 15,
"role": "CTA",
"narration": "旧口播1",
"visual": "旧画面1",
"speaker": None,
"product_exposure": "",
"entity_refs": [],
"dialogue": [],
},
],
}
base = ScriptVersion.objects.create(
project=self.project,
title="基准稿",
content=json.dumps(base_draft, ensure_ascii=False),
metadata={"hook": "原钩子", "entities": [], "total_duration": 30},
)
ScriptSegment.objects.create(
script_version=base,
sort_order=0,
duration_seconds=15,
narration="旧口播0",
visual_prompt="旧画面0",
)
ScriptSegment.objects.create(
script_version=base,
sort_order=1,
duration_seconds=15,
narration="旧口播1",
visual_prompt="旧画面1",
)
revised = json.loads(json.dumps(base_draft))
revised["segments"][1]["narration"] = "流式优化口播"
revised["segments"][1]["visual"] = "流式优化画面"
provider = Mock()
provider.chat_completion_stream.return_value = iter(
[
{"type": "delta", "text": json.dumps(revised, ensure_ascii=False)},
{"type": "done"},
]
)
self.provider_mocks[primary.id] = provider
events = self.parse_events(
list(
stream_script_agent(
project=self.project,
user=self.user,
model_config=primary,
mode="revise",
user_prompt="只改第二镜",
base_version_id=str(base.id),
aspect_ratio="9:16",
total_duration=60,
target_index=1,
)
)
)
task = AITask.objects.get(task_type=AITask.Type.SCRIPT_OPTIMIZATION)
saved = ScriptVersion.objects.filter(project=self.project).order_by("-created_at").first()
segments = list(saved.segments.order_by("sort_order"))
self.assertEqual(events[-1]["type"], "done")
self.assertEqual(len(segments), 2)
self.assertEqual(segments[0].narration, "旧口播0")
self.assertEqual(segments[1].narration, "流式优化口播")
attempt = task.model_attempts.get()
self.assertEqual(attempt.request_summary["target_index"], 1)
self.assertEqual(attempt.request_summary["total_duration"], 30)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
@@ -0,0 +1,696 @@
from decimal import Decimal
from io import BytesIO
from unittest.mock import Mock, patch
import requests
from django.test import TestCase
from apps.accounts.models import Team, TeamMember, User
from apps.ai.models import AIModelAttempt, AITask, ModelConfig, ModelProvider
from apps.ai.services import enqueue_standalone_images, run_standalone_image_task
from apps.assets.models import Asset, AssetFile
from apps.billing.models import CreditAccount, CreditLedger
from apps.products.models import Product
def image_metadata(
*,
outbound,
candidate=True,
ratios=None,
base_cost="0.50",
reference_modes=None,
max_reference_images=0,
):
modes = reference_modes or ["none"]
operations = ["image_generate"]
if set(modes) & {"single", "multiple"}:
operations.append("image_edit")
return {
"routing": {"fallback_on_failure": outbound, "fallback_candidate": candidate},
"capabilities": {
"operations": operations,
"features": [],
"reference_modes": modes,
"max_reference_images": max_reference_images,
"aspect_ratios": ratios or ["1:1", "3:4", "4:5", "9:16", "16:9"],
},
"pricing": {"base_cost_yuan": base_cost},
}
class StandaloneSingleImageRoutingTests(TestCase):
def setUp(self):
ModelConfig.objects.filter(capability=ModelConfig.Capability.IMAGE).update(status=ModelConfig.Status.DISABLED)
self.user = User.objects.create_user(username="standalone-routing", password="x")
self.team = Team.objects.create(name="Standalone Routing", owner=self.user)
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
CreditAccount.objects.create(team=self.team, balance=Decimal("1000"))
self.provider_mocks = {}
provider_patch = patch("apps.ai.services.get_image_provider", side_effect=self._provider_for)
provider_patch.start()
media_patch = patch(
"apps.ai.services.VolcanoArkProvider.media_to_bytes",
return_value=(BytesIO(b"image"), "image/png"),
)
self.media_to_bytes = media_patch.start()
storage_patch = patch("apps.ai.services.TosStorage")
storage = storage_patch.start()
stored = storage.return_value.upload_fileobj.return_value
stored.object_key = "generated.png"
stored.bucket = "bucket"
stored.content_type = "image/png"
stored.size_bytes = 5
self.addCleanup(patch.stopall)
def provider(self, name, priority):
provider, _ = ModelProvider.objects.update_or_create(
name=name,
defaults={
"display_name": name,
"status": ModelProvider.Status.ACTIVE,
"metadata": {"routing": {"fallback_priority": priority}},
},
)
return provider
def model(
self,
provider,
name,
*,
outbound=True,
ratios=None,
base_cost="0.50",
reference_modes=None,
max_reference_images=0,
):
return ModelConfig.objects.create(
provider=provider,
name=name,
display_name=name,
capability=ModelConfig.Capability.IMAGE,
endpoint="images/generations",
unit_price=Decimal("20"),
status=ModelConfig.Status.ACTIVE,
metadata=image_metadata(
outbound=outbound,
ratios=ratios,
base_cost=base_cost,
reference_modes=reference_modes,
max_reference_images=max_reference_images,
),
)
def _provider_for(self, model):
return self.provider_mocks.setdefault(model.id, self._new_provider_mock())
@staticmethod
def _new_provider_mock():
provider = Mock()
provider.image_generation.return_value = {"data": [{"url": "http://example.test/result.png"}]}
provider.image_edit.return_value = {"data": [{"url": "http://example.test/result.png"}]}
provider.extract_first_media_url.side_effect = lambda response: response["data"][0]["url"]
return provider
@staticmethod
def _new_generation_only_provider_mock():
provider = Mock(spec=["image_generation", "extract_first_media_url"])
provider.image_generation.return_value = {"data": [{"url": "http://example.test/result.png"}]}
provider.extract_first_media_url.side_effect = lambda response: response["data"][0]["url"]
return provider
def reference_asset(self, name):
asset = Asset.objects.create(
team=self.team,
created_by=self.user,
name=name,
asset_type=Asset.Type.IMAGE,
source=Asset.Source.UPLOAD,
category=Asset.Category.UPLOAD,
)
AssetFile.objects.create(
asset=asset,
object_key=f"{name}.png",
bucket="bucket",
content_type="image/png",
preview_url=f"http://example.test/{name}.png",
is_primary=True,
)
return asset
def submit(self, primary, *, count=1, ratio="1:1", reference_image_ids=None):
task = enqueue_standalone_images(
team=self.team,
user=self.user,
prompt="一只在太空漂浮的橙色猫",
mode="image",
count=count,
ratio=ratio,
image_model=f"{primary.provider.name}:{primary.name}",
reference_image_ids=reference_image_ids,
dispatch=False,
)
return task
def tryon_assets(self):
product = Product.objects.create(
team=self.team,
created_by=self.user,
title="橙色童装上衣",
category="服饰内衣",
)
cover = self.reference_asset("tryon-product")
cover.category = Asset.Category.PRODUCT_IMAGE
cover.save(update_fields=["category"])
product.cover_asset = cover
product.save(update_fields=["cover_asset"])
model_asset = self.reference_asset("tryon-model")
model_asset.source = Asset.Source.AI_GENERATED
model_asset.category = Asset.Category.PERSON
model_asset.save(update_fields=["source", "category"])
return product, model_asset
def submit_tryon(self, primary, *, count=1, ratio="4:5"):
product, model_asset = self.tryon_assets()
tasks = enqueue_standalone_images(
team=self.team,
user=self.user,
prompt="保持商品和模特身份,生成自然站姿上身图",
mode="model",
count=count,
product_id=str(product.id),
model_id=str(model_asset.id),
ratio=ratio,
image_model=f"{primary.provider.name}:{primary.name}",
dispatch=False,
)
return tasks
def platform_product(self):
product = Product.objects.create(
team=self.team,
created_by=self.user,
title="橙色童装上衣",
category="服饰内衣",
)
cover = self.reference_asset("platform-product")
cover.category = Asset.Category.PRODUCT_IMAGE
cover.save(update_fields=["category"])
product.cover_asset = cover
product.save(update_fields=["cover_asset"])
return product
def submit_platform(self, primary, *, count=1, ratio="4:5", platform_id="taobao"):
product = self.platform_product()
return enqueue_standalone_images(
team=self.team,
user=self.user,
prompt="保持商品一致,生成平台电商套图",
mode="cover",
count=count,
product_id=str(product.id),
ratio=ratio,
platform_id=platform_id,
image_model=f"{primary.provider.name}:{primary.name}",
dispatch=False,
)
def ledger_count(self, task, ledger_type):
return CreditLedger.objects.filter(task=task, ledger_type=ledger_type).count()
def test_user_selected_primary_is_first_and_success_never_queries_candidates(self):
primary = self.model(self.provider("single-primary", 100), "user-selected")
candidate = self.model(self.provider("single-unused", 1), "unused")
task = self.submit(primary)[0]
with patch("apps.ai.routing_executor.resolve_fallback_candidates") as resolver:
run_standalone_image_task(task_id=str(task.id))
resolver.assert_not_called()
task.refresh_from_db()
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
self.assertEqual(task.model_config_id, primary.id)
self.assertEqual(task.model_attempts.count(), 1)
attempt = task.model_attempts.get()
self.assertEqual(attempt.model_config_id, primary.id)
self.assertEqual(attempt.public_model_name, "AirShelf Image")
self.assertEqual(attempt.platform_cost, Decimal("0.5000"))
self.provider_mocks[primary.id].image_generation.assert_called_once()
self.assertNotIn(candidate.id, self.provider_mocks)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
def test_primary_retry_then_dynamic_candidate_success_has_one_charge(self):
primary = self.model(self.provider("single-fallback-primary", 100), "primary", base_cost="0.25")
candidate = self.model(self.provider("volcano", 10), "candidate", outbound=False, base_cost="0.75")
task = self.submit(primary, ratio="9:16")[0]
self.provider_mocks[primary.id] = self._new_provider_mock()
self.provider_mocks[primary.id].image_generation.side_effect = requests.ConnectionError("primary offline")
run_standalone_image_task(task_id=str(task.id))
task.refresh_from_db()
attempts = list(task.model_attempts.all())
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
self.assertEqual(task.model_config_id, primary.id)
self.assertEqual([item.model_config_id for item in attempts], [primary.id, primary.id, candidate.id])
self.assertEqual([item.is_retry for item in attempts], [False, True, False])
self.assertEqual([item.is_fallback for item in attempts], [False, False, True])
self.assertTrue(all(item.public_model_name == "AirShelf Image" for item in attempts))
self.assertEqual(task.base_cost, Decimal("0.7500"))
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
self.assertEqual(Asset.objects.filter(origin_task=task).count(), 1)
def test_direct_primary_switch_off_retries_but_does_not_call_candidate(self):
primary = self.model(self.provider("volcano", 10), "direct-primary", outbound=False)
candidate = self.model(self.provider("single-direct-unused", 20), "unused")
task = self.submit(primary)[0]
self.provider_mocks[primary.id] = self._new_provider_mock()
self.provider_mocks[primary.id].image_generation.side_effect = requests.ConnectionError("direct offline")
run_standalone_image_task(task_id=str(task.id))
task.refresh_from_db()
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(task.model_attempts.count(), 2)
self.assertNotIn(candidate.id, self.provider_mocks)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
def test_all_candidates_fail_releases_original_reservation_once(self):
primary = self.model(self.provider("single-all-primary", 100), "primary")
candidate = self.model(self.provider("volcano", 10), "candidate", outbound=False)
task = self.submit(primary)[0]
for model in (primary, candidate):
self.provider_mocks[model.id] = self._new_provider_mock()
self.provider_mocks[model.id].image_generation.side_effect = requests.ConnectionError("offline")
run_standalone_image_task(task_id=str(task.id))
task.refresh_from_db()
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(task.model_attempts.count(), 3)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
account = CreditAccount.objects.get(team=self.team)
self.assertEqual(account.reserved_balance, Decimal("0"))
def test_invalid_input_neither_retries_nor_fallbacks(self):
primary = self.model(self.provider("single-invalid-primary", 100), "primary")
candidate = self.model(self.provider("single-invalid-unused", 10), "unused")
task = self.submit(primary)[0]
self.provider_mocks[primary.id] = self._new_provider_mock()
self.provider_mocks[primary.id].image_generation.side_effect = RuntimeError("invalid image parameter")
run_standalone_image_task(task_id=str(task.id))
self.assertEqual(task.model_attempts.count(), 1)
self.assertEqual(task.model_attempts.get().error_type, "invalid_input")
self.assertNotIn(candidate.id, self.provider_mocks)
def test_incompatible_ratio_candidate_is_excluded(self):
primary = self.model(self.provider("single-ratio-primary", 100), "primary")
incompatible = self.model(
self.provider("single-ratio-incompatible", 10),
"incompatible",
ratios=["1:1"],
)
task = self.submit(primary, ratio="9:16")[0]
self.provider_mocks[primary.id] = self._new_provider_mock()
self.provider_mocks[primary.id].image_generation.side_effect = requests.ConnectionError("offline")
run_standalone_image_task(task_id=str(task.id))
task.refresh_from_db()
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(task.model_attempts.count(), 2)
self.assertNotIn(incompatible.id, self.provider_mocks)
def test_postprocess_failure_does_not_reinvoke_model_or_fallback(self):
primary = self.model(self.provider("single-post-primary", 100), "primary")
candidate = self.model(self.provider("single-post-unused", 10), "unused")
task = self.submit(primary)[0]
self.media_to_bytes.side_effect = RuntimeError("download failed")
run_standalone_image_task(task_id=str(task.id))
task.refresh_from_db()
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(task.model_attempts.count(), 1)
self.assertEqual(task.model_attempts.get().status, AIModelAttempt.Status.SUCCEEDED)
self.provider_mocks[primary.id].image_generation.assert_called_once()
self.assertNotIn(candidate.id, self.provider_mocks)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
def test_single_reference_uses_selected_edit_model_and_records_attempt(self):
primary = self.model(
self.provider("single-ref-primary", 20),
"selected-edit-model",
reference_modes=["none", "single", "multiple"],
max_reference_images=9,
)
reference = self.reference_asset("single-ref")
task = self.submit(primary, reference_image_ids=[str(reference.id)])[0]
run_standalone_image_task(task_id=str(task.id))
task.refresh_from_db()
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
self.assertEqual(task.model_config_id, primary.id)
self.assertEqual(task.model_attempts.count(), 1)
attempt = task.model_attempts.get()
self.assertEqual(attempt.operation, "image_edit")
self.assertEqual(attempt.request_summary["reference_images"], 1)
call = self.provider_mocks[primary.id].image_edit.call_args
self.assertEqual(call.kwargs["images"], ["http://example.test/single-ref.png"])
self.assertIn("唯一依据", call.kwargs["prompt"])
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
def test_selected_seedream_reference_model_stays_primary_and_uses_image_input(self):
primary = self.model(
self.provider("volcano", 10),
"seedream-selected",
outbound=False,
reference_modes=["none", "single", "multiple"],
max_reference_images=9,
)
reference = self.reference_asset("seedream-ref")
self.provider_mocks[primary.id] = self._new_generation_only_provider_mock()
task = self.submit(primary, reference_image_ids=[str(reference.id)])[0]
run_standalone_image_task(task_id=str(task.id))
task.refresh_from_db()
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
self.assertEqual(task.model_config_id, primary.id)
self.assertEqual(task.model_attempts.count(), 1)
call = self.provider_mocks[primary.id].image_generation.call_args
self.assertEqual(call.kwargs["image"], ["http://example.test/seedream-ref.png"])
self.assertEqual(task.model_attempts.get().public_model_name, primary.display_name)
def test_multiple_references_retry_then_dynamic_compatible_candidate_success(self):
primary = self.model(
self.provider("multi-ref-primary", 100),
"primary-edit",
reference_modes=["none", "single", "multiple"],
max_reference_images=9,
)
candidate = self.model(
self.provider("volcano", 10),
"candidate-edit",
outbound=False,
reference_modes=["none", "single", "multiple"],
max_reference_images=9,
)
references = [self.reference_asset("multi-ref-1"), self.reference_asset("multi-ref-2")]
task = self.submit(primary, reference_image_ids=[str(item.id) for item in references])[0]
self.provider_mocks[primary.id] = self._new_provider_mock()
self.provider_mocks[primary.id].image_edit.side_effect = requests.ConnectionError("primary offline")
run_standalone_image_task(task_id=str(task.id))
task.refresh_from_db()
attempts = list(task.model_attempts.all())
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
self.assertEqual([item.model_config_id for item in attempts], [primary.id, primary.id, candidate.id])
self.assertTrue(attempts[-1].is_fallback)
self.assertTrue(all(item.operation == "image_edit" for item in attempts))
candidate_call = self.provider_mocks[candidate.id].image_edit.call_args
self.assertEqual(
candidate_call.kwargs["images"],
["http://example.test/multi-ref-1.png", "http://example.test/multi-ref-2.png"],
)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
def test_multiple_references_exclude_candidate_with_insufficient_limit(self):
primary = self.model(
self.provider("multi-limit-primary", 100),
"primary-edit",
reference_modes=["none", "single", "multiple"],
max_reference_images=9,
)
incompatible = self.model(
self.provider("multi-limit-incompatible", 10),
"single-only",
reference_modes=["none", "single", "multiple"],
max_reference_images=1,
)
references = [self.reference_asset("limit-ref-1"), self.reference_asset("limit-ref-2")]
task = self.submit(primary, reference_image_ids=[str(item.id) for item in references])[0]
self.provider_mocks[primary.id] = self._new_provider_mock()
self.provider_mocks[primary.id].image_edit.side_effect = requests.ConnectionError("primary offline")
run_standalone_image_task(task_id=str(task.id))
task.refresh_from_db()
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(task.model_attempts.count(), 2)
self.assertNotIn(incompatible.id, self.provider_mocks)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
def test_tryon_selected_edit_model_is_primary_and_records_two_references(self):
primary = self.model(
self.provider("tryon-primary", 20),
"selected-tryon-edit",
reference_modes=["none", "single", "multiple"],
max_reference_images=9,
)
task = self.submit_tryon(primary)[0]
run_standalone_image_task(task_id=str(task.id))
task.refresh_from_db()
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
self.assertEqual(task.model_config_id, primary.id)
self.assertEqual(task.model_attempts.count(), 1)
attempt = task.model_attempts.get()
self.assertEqual(attempt.operation, "image_edit")
self.assertEqual(attempt.request_summary["reference_images"], 2)
call = self.provider_mocks[primary.id].image_edit.call_args
self.assertEqual(
call.kwargs["images"],
["http://example.test/tryon-product.png", "http://example.test/tryon-model.png"],
)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
def test_tryon_selected_seedream_stays_primary_and_uses_image_generation(self):
primary = self.model(
self.provider("volcano", 10),
"seedream-tryon-selected",
outbound=False,
reference_modes=["none", "single", "multiple"],
max_reference_images=9,
)
self.provider_mocks[primary.id] = self._new_generation_only_provider_mock()
task = self.submit_tryon(primary)[0]
run_standalone_image_task(task_id=str(task.id))
task.refresh_from_db()
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
self.assertEqual(task.model_config_id, primary.id)
self.assertEqual(task.model_attempts.count(), 1)
call = self.provider_mocks[primary.id].image_generation.call_args
self.assertEqual(len(call.kwargs["image"]), 2)
self.assertEqual(task.model_attempts.get().public_model_name, primary.display_name)
def test_tryon_primary_retry_then_dynamic_candidate_success_has_one_charge(self):
primary = self.model(
self.provider("tryon-fallback-primary", 100),
"primary-tryon",
reference_modes=["none", "single", "multiple"],
max_reference_images=9,
)
candidate = self.model(
self.provider("volcano", 10),
"candidate-tryon",
outbound=False,
reference_modes=["none", "single", "multiple"],
max_reference_images=9,
)
task = self.submit_tryon(primary)[0]
self.provider_mocks[primary.id] = self._new_provider_mock()
self.provider_mocks[primary.id].image_edit.side_effect = requests.ConnectionError("primary offline")
run_standalone_image_task(task_id=str(task.id))
task.refresh_from_db()
attempts = list(task.model_attempts.all())
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
self.assertEqual([item.model_config_id for item in attempts], [primary.id, primary.id, candidate.id])
self.assertTrue(attempts[-1].is_fallback)
self.assertTrue(all(item.request_summary["reference_images"] == 2 for item in attempts))
self.assertTrue(all(item.public_model_name == "AirShelf Image" for item in attempts))
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
def test_tryon_direct_primary_failure_retries_without_fallback_and_releases(self):
primary = self.model(
self.provider("volcano", 10),
"direct-tryon",
outbound=False,
reference_modes=["none", "single", "multiple"],
max_reference_images=9,
)
candidate = self.model(
self.provider("tryon-direct-unused", 20),
"unused-tryon",
reference_modes=["none", "single", "multiple"],
max_reference_images=9,
)
self.provider_mocks[primary.id] = self._new_generation_only_provider_mock()
self.provider_mocks[primary.id].image_generation.side_effect = requests.ConnectionError("direct offline")
task = self.submit_tryon(primary)[0]
run_standalone_image_task(task_id=str(task.id))
task.refresh_from_db()
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(task.model_attempts.count(), 2)
self.assertNotIn(candidate.id, self.provider_mocks)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
def test_tryon_batch_keeps_independent_task_attempt_and_billing_chains(self):
primary = self.model(
self.provider("tryon-batch-primary", 20),
"batch-tryon",
reference_modes=["none", "single", "multiple"],
max_reference_images=9,
)
tasks = self.submit_tryon(primary, count=2)
for task in tasks:
run_standalone_image_task(task_id=str(task.id))
for task in tasks:
task.refresh_from_db()
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
self.assertEqual(task.model_attempts.count(), 1)
self.assertEqual(task.model_attempts.get().request_summary["reference_images"], 2)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
def test_platform_kit_selected_model_is_primary_and_keeps_platform_prompt(self):
primary = self.model(
self.provider("platform-primary", 20),
"selected-platform-edit",
reference_modes=["none", "single", "multiple"],
max_reference_images=9,
)
task = self.submit_platform(primary, platform_id="taobao")[0]
run_standalone_image_task(task_id=str(task.id))
task.refresh_from_db()
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
self.assertEqual(task.model_config_id, primary.id)
self.assertEqual(task.request_payload["platform_id"], "taobao")
self.assertEqual(task.request_payload["platform_name"], "淘宝")
attempt = task.model_attempts.get()
self.assertEqual(attempt.operation, "image_edit")
self.assertEqual(attempt.request_summary["reference_images"], 1)
call = self.provider_mocks[primary.id].image_edit.call_args
self.assertEqual(call.kwargs["images"], ["http://example.test/platform-product.png"])
self.assertIn("淘宝", call.kwargs["prompt"])
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
def test_platform_kit_selected_seedream_stays_primary(self):
primary = self.model(
self.provider("volcano", 10),
"seedream-platform-selected",
outbound=False,
reference_modes=["none", "single", "multiple"],
max_reference_images=9,
)
self.provider_mocks[primary.id] = self._new_generation_only_provider_mock()
task = self.submit_platform(primary, platform_id="douyin")[0]
run_standalone_image_task(task_id=str(task.id))
task.refresh_from_db()
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
self.assertEqual(task.model_config_id, primary.id)
self.assertEqual(task.model_attempts.count(), 1)
call = self.provider_mocks[primary.id].image_generation.call_args
self.assertEqual(call.kwargs["image"], ["http://example.test/platform-product.png"])
self.assertEqual(task.model_attempts.get().public_model_name, primary.display_name)
def test_platform_kit_retry_then_dynamic_candidate_success_has_one_charge(self):
primary = self.model(
self.provider("platform-fallback-primary", 100),
"primary-platform",
reference_modes=["none", "single", "multiple"],
max_reference_images=9,
)
candidate = self.model(
self.provider("volcano", 10),
"candidate-platform",
outbound=False,
reference_modes=["none", "single", "multiple"],
max_reference_images=9,
)
task = self.submit_platform(primary)[0]
self.provider_mocks[primary.id] = self._new_provider_mock()
self.provider_mocks[primary.id].image_edit.side_effect = requests.ConnectionError("primary offline")
run_standalone_image_task(task_id=str(task.id))
task.refresh_from_db()
attempts = list(task.model_attempts.all())
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
self.assertEqual([item.model_config_id for item in attempts], [primary.id, primary.id, candidate.id])
self.assertTrue(attempts[-1].is_fallback)
self.assertTrue(all(item.request_summary["reference_images"] == 1 for item in attempts))
self.assertTrue(all(item.public_model_name == "AirShelf Image" for item in attempts))
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
def test_platform_kit_batch_keeps_independent_tasks_attempts_and_billing(self):
primary = self.model(
self.provider("platform-batch-primary", 20),
"batch-platform",
reference_modes=["none", "single", "multiple"],
max_reference_images=9,
)
tasks = self.submit_platform(primary, count=3, platform_id="jd")
for task in tasks:
run_standalone_image_task(task_id=str(task.id))
self.assertEqual({task.request_payload["batch_id"] for task in tasks}, {tasks[0].request_payload["batch_id"]})
for index, task in enumerate(tasks):
task.refresh_from_db()
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
self.assertEqual(task.request_payload["index"], index)
self.assertEqual(task.request_payload["platform_id"], "jd")
self.assertEqual(task.model_attempts.count(), 1)
self.assertEqual(task.model_attempts.get().request_summary["reference_images"], 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
def test_batch_without_references_keeps_legacy_path_until_its_own_step(self):
primary = self.model(self.provider("single-batch-primary", 100), "primary")
tasks = self.submit(primary, count=2)
for task in tasks:
run_standalone_image_task(task_id=str(task.id))
self.assertTrue(all("model_routing_v1" not in task.request_payload for task in tasks))
self.assertEqual(AIModelAttempt.objects.filter(task__in=tasks).count(), 0)
self.assertEqual(self.provider_mocks[primary.id].image_generation.call_count, 2)
@@ -0,0 +1,312 @@
from decimal import Decimal
from unittest.mock import Mock, patch
import requests
from django.test import TestCase, override_settings
from apps.accounts.models import Team, TeamMember, User
from apps.ai.models import AITask, ModelConfig, ModelProvider
from apps.ai.services import _storyboard_shot_worker, poll_storyboard
from apps.assets.models import Asset
from apps.billing.models import CreditAccount, CreditLedger
from apps.products.models import Product
from apps.projects.models import (
Project,
ScriptSegment,
ScriptVersion,
StoryboardShot,
StoryboardShotVersion,
)
def _metadata(*, outbound=True, base_cost="0.50", max_refs=9):
return {
"routing": {"fallback_on_failure": outbound, "fallback_candidate": True},
"capabilities": {
"operations": ["image_generate", "image_edit"],
"features": [],
"reference_modes": ["none", "single", "multiple"],
"max_reference_images": max_refs,
"aspect_ratios": ["1:1", "3:4", "4:5", "9:16", "16:9"],
},
"pricing": {"base_cost_yuan": base_cost},
}
@override_settings(STORYBOARD_MAX_PARALLEL=4)
class StoryboardRoutingTests(TestCase):
def setUp(self):
ModelConfig.objects.filter(capability=ModelConfig.Capability.IMAGE).update(
status=ModelConfig.Status.DISABLED
)
self.user = User.objects.create_user(username="storyboard-routing", password="x")
self.team = Team.objects.create(name="Storyboard Routing", owner=self.user)
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
CreditAccount.objects.create(team=self.team, balance=Decimal("1000"))
product = Product.objects.create(team=self.team, created_by=self.user, title="测试商品")
self.project = Project.objects.create(
team=self.team,
created_by=self.user,
product=product,
name="故事板路由项目",
)
script = ScriptVersion.objects.create(
project=self.project,
title="脚本",
content="结构化脚本",
is_adopted=True,
)
self.segment = ScriptSegment.objects.create(
script_version=script,
sort_order=0,
visual_prompt="人物在商品旁展示",
)
self.shot = StoryboardShot.objects.create(
project=self.project,
script_segment=self.segment,
sort_order=0,
status=StoryboardShot.Status.QUEUED,
)
self.refs = [
{"url": "http://example.test/person.png", "label": "女主", "type": "person"},
{"url": "http://example.test/product.png", "label": "商品", "type": "product"},
]
self.provider_mocks = {}
patch("apps.ai.services.get_image_provider", side_effect=self._provider_for).start()
patch(
"apps.ai.services._storyboard_reference_images",
side_effect=lambda project, segment: list(self.refs),
).start()
patch(
"apps.ai.services.build_storyboard_frame_prompt_refs",
return_value="带参考图编号与锁定约束的故事板提示词",
).start()
patch(
"apps.ai.services.build_storyboard_frame_prompt",
return_value="无参考图故事板提示词",
).start()
patch("apps.ai.services._store_generated_media", side_effect=self._store_media).start()
patch("apps.assets.review.submit_asset_for_review").start()
patch("apps.ai.services.notify_generation_failure").start()
self.addCleanup(patch.stopall)
def provider(self, name, priority):
return ModelProvider.objects.create(
name=name,
display_name=name,
status=ModelProvider.Status.ACTIVE,
metadata={"routing": {"fallback_priority": priority}},
)
def model(self, provider, name, *, outbound=True, base_cost="0.50", max_refs=9):
return ModelConfig.objects.create(
provider=provider,
name=name,
display_name=name,
capability=ModelConfig.Capability.IMAGE,
endpoint="images/generations",
unit_price=Decimal("20"),
status=ModelConfig.Status.ACTIVE,
metadata=_metadata(outbound=outbound, base_cost=base_cost, max_refs=max_refs),
)
@staticmethod
def _new_provider_mock():
provider = Mock()
response = {"data": [{"url": "http://example.test/storyboard.png"}]}
provider.image_edit.return_value = response
provider.image_generation.return_value = response
provider.extract_first_media_url.side_effect = lambda value: value["data"][0]["url"]
return provider
@staticmethod
def _new_generation_only_provider_mock():
provider = Mock(spec=["image_generation", "extract_first_media_url"])
provider.image_generation.return_value = {
"data": [{"url": "http://example.test/storyboard.png"}]
}
provider.extract_first_media_url.side_effect = lambda value: value["data"][0]["url"]
return provider
def _provider_for(self, model):
return self.provider_mocks.setdefault(model.id, self._new_provider_mock())
def _store_media(self, **kwargs):
return Asset.objects.create(
team=kwargs["team"],
created_by=kwargs["user"],
name=kwargs["name"],
asset_type=kwargs["asset_type"],
source=Asset.Source.AI_GENERATED,
category=kwargs["category"],
origin_task=kwargs["task"],
)
def enqueue(self, *, run=True):
with patch("threading.Thread"):
result = poll_storyboard(project=self.project, user=self.user)
self.assertEqual(result["status"], "generating")
task = AITask.objects.filter(project=self.project, task_type=AITask.Type.STORYBOARD).latest(
"created_at"
)
if run:
_storyboard_shot_worker(str(task.id), str(self.shot.id), str(self.user.id))
return task
@staticmethod
def ledger_count(task, ledger_type):
return CreditLedger.objects.filter(task=task, ledger_type=ledger_type).count()
def test_multi_reference_success_records_9_16_and_adopts_one_shot_version(self):
primary = self.model(self.provider("storyboard-primary", 20), "storyboard-primary")
task = self.enqueue()
task.refresh_from_db()
self.shot.refresh_from_db()
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
self.assertTrue(task.request_payload["model_routing_v1"])
self.assertEqual(task.base_cost, Decimal("0.5000"))
attempt = task.model_attempts.get()
self.assertEqual(attempt.model_config_id, primary.id)
self.assertEqual(attempt.operation, "image_edit")
self.assertEqual(attempt.request_summary["reference_images"], 2)
self.assertEqual(attempt.request_summary["aspect_ratio"], "9:16")
self.assertEqual(attempt.request_summary["storyboard_shot"], str(self.shot.id))
call = self.provider_mocks[primary.id].image_edit.call_args
self.assertEqual(
call.kwargs["images"],
["http://example.test/person.png", "http://example.test/product.png"],
)
self.assertEqual(call.kwargs["prompt"], "带参考图编号与锁定约束的故事板提示词")
self.assertEqual(call.kwargs["size"], "1024x1536")
self.assertIsNotNone(self.shot.adopted_version_id)
self.assertEqual(self.shot.adopted_version.task_id, task.id)
self.assertEqual(self.shot.versions.count(), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
def test_retry_then_dynamic_candidate_success_charges_once(self):
primary = self.model(
self.provider("storyboard-fallback-primary", 100),
"storyboard-primary",
base_cost="0.25",
)
candidate = self.model(
self.provider("volcano", 10),
"storyboard-candidate",
outbound=False,
base_cost="0.75",
)
self.provider_mocks[primary.id] = self._new_provider_mock()
self.provider_mocks[primary.id].image_edit.side_effect = requests.ConnectionError("offline")
task = self.enqueue()
task.refresh_from_db()
attempts = list(task.model_attempts.all())
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
self.assertEqual([a.model_config_id for a in attempts], [primary.id, primary.id, candidate.id])
self.assertTrue(attempts[-1].is_fallback)
self.assertEqual(task.base_cost, Decimal("0.7500"))
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
def test_all_candidates_fail_releases_and_keeps_old_adopted_version(self):
primary = self.model(self.provider("storyboard-all-primary", 100), "storyboard-primary")
candidate = self.model(
self.provider("volcano", 10), "storyboard-candidate", outbound=False
)
old_asset = Asset.objects.create(
team=self.team,
created_by=self.user,
name="旧故事板",
asset_type=Asset.Type.IMAGE,
source=Asset.Source.AI_GENERATED,
category=Asset.Category.STORYBOARD,
)
old_version = StoryboardShotVersion.objects.create(
shot=self.shot,
asset=old_asset,
prompt="旧画面",
is_adopted=True,
)
self.shot.adopted_version = old_version
self.shot.save(update_fields=["adopted_version", "updated_at"])
for model in (primary, candidate):
self.provider_mocks[model.id] = self._new_provider_mock()
self.provider_mocks[model.id].image_edit.side_effect = requests.ConnectionError("offline")
task = self.enqueue()
task.refresh_from_db()
self.shot.refresh_from_db()
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(task.model_attempts.count(), 3)
self.assertEqual(self.shot.adopted_version_id, old_version.id)
self.assertEqual(self.shot.versions.count(), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
def test_no_reference_uses_image_generation_and_keeps_vertical_capability(self):
primary = self.model(self.provider("storyboard-no-ref", 20), "storyboard-no-ref")
self.refs = []
task = self.enqueue()
task.refresh_from_db()
attempt = task.model_attempts.get()
self.assertEqual(attempt.operation, "image_generate")
self.assertEqual(attempt.request_summary["reference_images"], 0)
self.assertEqual(attempt.request_summary["aspect_ratio"], "9:16")
call = self.provider_mocks[primary.id].image_generation.call_args
self.assertEqual(call.kwargs["prompt"], "无参考图故事板提示词")
self.assertNotIn("size", call.kwargs)
def test_direct_primary_uses_image_generation_with_all_references(self):
primary = self.model(
self.provider("volcano", 10), "seedream-storyboard", outbound=False
)
self.provider_mocks[primary.id] = self._new_generation_only_provider_mock()
task = self.enqueue()
task.refresh_from_db()
call = self.provider_mocks[primary.id].image_generation.call_args
self.assertEqual(
call.kwargs["image"],
["http://example.test/person.png", "http://example.test/product.png"],
)
self.assertEqual(call.kwargs["size"], "1024x1536")
self.assertEqual(task.model_attempts.get().public_model_name, primary.display_name)
def test_poll_creates_one_independent_task_and_reservation_per_shot(self):
self.model(self.provider("storyboard-batch", 20), "storyboard-batch")
second_segment = ScriptSegment.objects.create(
script_version=self.segment.script_version,
sort_order=1,
visual_prompt="第二镜",
)
second_shot = StoryboardShot.objects.create(
project=self.project,
script_segment=second_segment,
sort_order=1,
status=StoryboardShot.Status.QUEUED,
)
with patch("threading.Thread"):
result = poll_storyboard(project=self.project, user=self.user)
self.assertEqual(result, {"status": "generating", "done": 0, "total": 2})
tasks = list(
AITask.objects.filter(project=self.project, task_type=AITask.Type.STORYBOARD).order_by(
"created_at"
)
)
self.assertEqual(len(tasks), 2)
self.assertEqual(
{task.request_payload["storyboard_shot"] for task in tasks},
{str(self.shot.id), str(second_shot.id)},
)
self.assertTrue(all(task.request_payload["model_routing_v1"] for task in tasks))
self.assertTrue(all(task.base_cost == Decimal("0") for task in tasks))
self.assertTrue(all(self.ledger_count(task, CreditLedger.Type.RESERVE) == 1 for task in tasks))
@@ -0,0 +1,278 @@
from copy import deepcopy
from datetime import timedelta
from decimal import Decimal
from unittest.mock import Mock, patch
import requests
from django.conf import settings
from django.test import TestCase, override_settings
from django.utils import timezone
from apps.accounts.models import Team, TeamMember, User
from apps.ai.models import AITask, ModelConfig, ModelProvider
from apps.ai.services import poll_video_segment, submit_video_segment
from apps.assets.models import Asset
from apps.billing.models import CreditAccount, CreditLedger
from apps.products.models import Product
from apps.projects.models import Project, VideoSegment
def _fast_routing_policy():
policy = deepcopy(settings.MODEL_ROUTING_POLICY)
policy["jitter_ratio"] = 0
policy["video"]["submit_retry_delays"] = [0]
return policy
def _video_metadata(*, outbound, candidate=True, price=46):
return {
"routing": {
"fallback_on_failure": outbound,
"fallback_candidate": candidate,
},
"capabilities": {
"operations": ["video_generate"],
"features": ["generate_audio"],
"max_reference_images": 9,
"max_reference_videos": 3,
"max_reference_audios": 3,
"aspect_ratios": ["9:16"],
"resolutions": ["720p"],
"durations": [15],
},
"pricing": {
"unit": "cny_per_million_tokens",
"default": {"no_ref_video": price, "with_ref_video": price},
},
}
@override_settings(MODEL_ROUTING_POLICY=_fast_routing_policy())
class VideoSegmentRoutingTests(TestCase):
def setUp(self):
ModelConfig.objects.filter(capability=ModelConfig.Capability.VIDEO).update(
status=ModelConfig.Status.DISABLED,
is_default=False,
)
self.user = User.objects.create_user(username="video-routing", password="x")
self.team = Team.objects.create(name="Video Routing", owner=self.user)
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
CreditAccount.objects.create(team=self.team, balance=Decimal("10000"))
product = Product.objects.create(team=self.team, created_by=self.user, title="Product")
self.project = Project.objects.create(
team=self.team,
created_by=self.user,
product=product,
name="Video routing project",
)
self.segment = VideoSegment.objects.create(
project=self.project,
sort_order=0,
target_duration_seconds=15,
)
self.provider_mocks = {}
patch("apps.ai.services.get_video_provider", side_effect=self._provider_for).start()
patch("apps.ai.services._video_reference_images", return_value=[]).start()
patch("apps.ai.services.build_video_segment_prompt", return_value="video prompt").start()
patch("apps.ai.services.notify_generation_failure").start()
self.addCleanup(patch.stopall)
def provider(self, name, priority):
return ModelProvider.objects.create(
name=name,
display_name=name,
status=ModelProvider.Status.ACTIVE,
base_url="https://video.example/v1",
metadata={"routing": {"fallback_priority": priority}},
)
def model(self, provider, name, *, outbound, price=46, default=False):
return ModelConfig.objects.create(
provider=provider,
name=name,
display_name=name,
capability=ModelConfig.Capability.VIDEO,
endpoint="contents/generations/tasks",
unit_price=Decimal("1"),
status=ModelConfig.Status.ACTIVE,
is_default=default,
metadata=_video_metadata(outbound=outbound, price=price),
)
def _provider_for(self, model):
return self.provider_mocks.setdefault(model.id, Mock())
def _submit(self):
submit_video_segment(video_segment=self.segment, user=self.user, prompt="test")
return AITask.objects.get(task_type=AITask.Type.VIDEO_SEGMENT, team=self.team)
@staticmethod
def ledger_count(task, ledger_type):
return CreditLedger.objects.filter(task=task, ledger_type=ledger_type).count()
@staticmethod
def _asset(task):
return Asset.objects.create(
team=task.team,
created_by=task.created_by,
name="video clip",
asset_type=Asset.Type.VIDEO,
source=Asset.Source.AI_GENERATED,
category=Asset.Category.VIDEO_CLIP,
origin_task=task,
)
def test_first_submit_success_creates_one_attempt_and_one_reservation(self):
primary = self.model(self.provider("video-primary-ok", 20), "video-primary-ok", outbound=True, default=True)
provider = self._provider_for(primary)
provider.create_video_task.return_value = {"id": "remote-primary", "status": "queued"}
task = self._submit()
task.refresh_from_db()
self.segment.refresh_from_db()
self.assertEqual(task.status, AITask.Status.SUBMITTED)
self.assertEqual(task.provider_task_id, "remote-primary")
self.assertEqual(task.request_payload["actual_model_config_id"], str(primary.id))
self.assertEqual(task.model_attempts.count(), 1)
self.assertEqual(task.model_attempts.get().provider_task_id, "remote-primary")
self.assertEqual(self.segment.status, VideoSegment.Status.RUNNING)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
@patch("apps.ai.services._store_generated_media")
def test_retry_then_dynamic_fallback_polls_actual_model_and_charges_once(self, store_media):
primary = self.model(self.provider("video-primary-fail", 20), "video-primary-fail", outbound=True, price=46, default=True)
candidate = self.model(self.provider("video-candidate-ok", 30), "video-candidate-ok", outbound=False, price=23)
primary_provider = self._provider_for(primary)
fallback_provider = self._provider_for(candidate)
primary_provider.create_video_task.side_effect = requests.ConnectionError("primary offline")
fallback_provider.create_video_task.return_value = {"id": "remote-fallback", "status": "queued"}
task = self._submit()
store_media.return_value = self._asset(task)
fallback_provider.poll_video_task.return_value = {
"status": "succeeded",
"usage": {"total_tokens": 300000},
"content": {"video_url": "https://video.example/result.mp4"},
}
fallback_provider.extract_first_media_url.return_value = "https://video.example/result.mp4"
version = poll_video_segment(video_segment=self.segment, user=self.user)
task.refresh_from_db()
attempts = list(task.model_attempts.order_by("sequence"))
self.assertIsNotNone(version)
self.assertEqual(
[attempt.model_config_id for attempt in attempts],
[primary.id, primary.id, candidate.id],
)
self.assertTrue(attempts[-1].is_fallback)
self.assertEqual(task.model_config_id, primary.id)
self.assertEqual(task.request_payload["actual_model_config_id"], str(candidate.id))
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
self.assertFalse(primary_provider.poll_video_task.called)
fallback_provider.poll_video_task.assert_called_once_with(
endpoint=candidate.endpoint,
provider_task_id="remote-fallback",
timeout=60.0,
)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
def test_direct_model_with_outbound_disabled_never_uses_candidate(self):
primary = self.model(self.provider("ark", 10), "video-direct-locked", outbound=False, default=True)
candidate = self.model(self.provider("video-unused", 20), "video-unused", outbound=False)
primary_provider = self._provider_for(primary)
unused_provider = self._provider_for(candidate)
primary_provider.create_video_task.side_effect = requests.ConnectionError("direct offline")
unused_provider.create_video_task.return_value = {"id": "must-not-run"}
with self.assertRaises(requests.ConnectionError):
self._submit()
task = AITask.objects.get(task_type=AITask.Type.VIDEO_SEGMENT, team=self.team)
self.assertEqual(
list(task.model_attempts.values_list("model_config_id", flat=True)),
[primary.id, primary.id],
)
self.assertFalse(unused_provider.create_video_task.called)
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
def test_submit_read_timeout_is_state_unknown_and_is_never_retried(self):
primary = self.model(self.provider("video-timeout", 20), "video-timeout", outbound=True, default=True)
candidate = self.model(self.provider("video-timeout-unused", 30), "video-timeout-unused", outbound=False)
primary_provider = self._provider_for(primary)
unused_provider = self._provider_for(candidate)
primary_provider.create_video_task.side_effect = requests.ReadTimeout("response lost")
with self.assertRaisesRegex(RuntimeError, "state unknown|状态未知"):
self._submit()
task = AITask.objects.get(task_type=AITask.Type.VIDEO_SEGMENT, team=self.team)
self.assertEqual(task.model_attempts.count(), 1)
self.assertFalse(unused_provider.create_video_task.called)
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
def test_submit_response_without_remote_id_is_never_retried(self):
primary = self.model(self.provider("video-no-id", 20), "video-no-id", outbound=True, default=True)
candidate = self.model(self.provider("video-no-id-unused", 30), "video-no-id-unused", outbound=False)
primary_provider = self._provider_for(primary)
unused_provider = self._provider_for(candidate)
primary_provider.create_video_task.return_value = {"status": "queued"}
with self.assertRaisesRegex(RuntimeError, "state unknown|状态未知"):
self._submit()
task = AITask.objects.get(task_type=AITask.Type.VIDEO_SEGMENT, team=self.team)
attempt = task.model_attempts.get()
self.assertTrue(attempt.response_summary["provider_task_id_missing"])
self.assertFalse(unused_provider.create_video_task.called)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
def test_poll_network_error_keeps_task_and_reservation_for_later_poll(self):
primary = self.model(self.provider("video-poll-unknown", 20), "video-poll-unknown", outbound=True, default=True)
provider = self._provider_for(primary)
provider.create_video_task.return_value = {"id": "remote-poll", "status": "queued"}
task = self._submit()
provider.poll_video_task.side_effect = requests.ConnectionError("poll offline")
with self.assertRaises(requests.ConnectionError):
poll_video_segment(video_segment=self.segment, user=self.user)
task.refresh_from_db()
self.segment.refresh_from_db()
self.assertEqual(task.status, AITask.Status.SUBMITTED)
self.assertEqual(self.segment.status, VideoSegment.Status.RUNNING)
self.assertEqual(task.model_attempts.count(), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
self.assertGreater(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
def test_generation_timeout_releases_without_polling_or_resubmitting(self):
primary = self.model(self.provider("video-generation-timeout", 20), "video-generation-timeout", outbound=True, default=True)
provider = self._provider_for(primary)
provider.create_video_task.return_value = {"id": "remote-slow", "status": "queued"}
task = self._submit()
task.submitted_at = timezone.now() - timedelta(seconds=1900)
task.save(update_fields=["submitted_at", "updated_at"])
result = poll_video_segment(video_segment=self.segment, user=self.user)
task.refresh_from_db()
self.segment.refresh_from_db()
self.assertIsNone(result)
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(self.segment.status, VideoSegment.Status.FAILED)
self.assertFalse(provider.poll_video_task.called)
self.assertEqual(provider.create_video_task.call_count, 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
@@ -0,0 +1,241 @@
from decimal import Decimal
from unittest.mock import Mock, patch
import requests
from django.test import SimpleTestCase, TestCase
from apps.accounts.models import Team, TeamMember, User
from apps.ai.models import AITask, ModelConfig, ModelProvider
from apps.ai.providers.openai_compatible import OpenAICompatibleProvider
from apps.ai.services import DEFAULT_VOICEOVER_VOICE, synthesize_project_voiceover
from apps.assets.models import Asset
from apps.assets.storage import StoredObject
from apps.billing.models import CreditAccount, CreditLedger
from apps.products.models import Product
from apps.projects.models import Project
def _metadata(*, outbound, candidate=True, voice="BV700_streaming", base_cost="0.10"):
return {
"routing": {
"fallback_on_failure": outbound,
"fallback_candidate": candidate,
},
"capabilities": {
"operations": ["tts"],
"features": [],
"languages": ["zh-CN"],
"voice_map": {DEFAULT_VOICEOVER_VOICE: voice},
"max_chars": 10000,
"speed_range": [0.5, 2.0],
"output_formats": ["mp3"],
},
"pricing": {
"chars_per_unit": 500,
"points_per_unit": 10,
"min_units": 1,
"base_cost_yuan_per_unit": base_cost,
},
}
class VoiceoverRoutingTests(TestCase):
def setUp(self):
ModelConfig.objects.filter(capability=ModelConfig.Capability.AUDIO).update(
status=ModelConfig.Status.DISABLED
)
self.user = User.objects.create_user(username="voice-routing", password="x")
self.team = Team.objects.create(name="Voice Routing", owner=self.user)
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
CreditAccount.objects.create(team=self.team, balance=Decimal("1000"))
product = Product.objects.create(team=self.team, created_by=self.user, title="测试商品")
self.project = Project.objects.create(
team=self.team,
created_by=self.user,
product=product,
name="配音路由项目",
)
self.direct_provider, _ = ModelProvider.objects.get_or_create(
name="ark",
defaults={"display_name": "ARK"},
)
self.direct_provider.status = ModelProvider.Status.ACTIVE
self.direct_provider.metadata = {"routing": {"fallback_priority": 10}}
self.direct_provider.save(update_fields=["status", "metadata", "updated_at"])
self.provider_mocks = {}
self.direct_mock = Mock()
self.direct_mock.configured = True
self.direct_mock.synthesize.return_value = (b"direct-mp3", 1200)
patch("apps.ai.services.get_audio_provider", side_effect=self._provider_for).start()
stored = StoredObject(
object_key="qa-voice.mp3",
bucket="qa",
content_type="audio/mpeg",
size_bytes=10,
)
patch("apps.ai.services.TosStorage.upload_fileobj", return_value=stored).start()
self.addCleanup(patch.stopall)
def provider(self, name, priority):
return ModelProvider.objects.create(
name=name,
display_name=name,
status=ModelProvider.Status.ACTIVE,
base_url="https://voice.example/v1",
metadata={"routing": {"fallback_priority": priority}},
)
def model(self, provider, name, *, outbound, voice="BV700_streaming", base_cost="0.10"):
return ModelConfig.objects.create(
provider=provider,
name=name,
display_name=name,
capability=ModelConfig.Capability.AUDIO,
endpoint="audio/speech",
unit_price=Decimal("10"),
status=ModelConfig.Status.ACTIVE,
metadata=_metadata(outbound=outbound, voice=voice, base_cost=base_cost),
)
def _provider_for(self, model):
if model.provider_id == self.direct_provider.id:
return self.direct_mock
return self.provider_mocks.setdefault(model.id, Mock())
@staticmethod
def ledger_count(task, ledger_type):
return CreditLedger.objects.filter(task=task, ledger_type=ledger_type).count()
def synthesize(self, items):
return synthesize_project_voiceover(
project=self.project,
user=self.user,
items=items,
voice_type=DEFAULT_VOICEOVER_VOICE,
speed_ratio=1.0,
)
def test_two_cues_create_two_attempts_but_one_reservation_and_charge(self):
primary = self.model(self.direct_provider, "voice-direct-success", outbound=False)
voiceover = self.synthesize(
[{"index": 0, "text": "第一镜旁白"}, {"index": 1, "text": "第二镜旁白"}]
)
task = AITask.objects.get(task_type=AITask.Type.VOICEOVER)
attempts = list(task.model_attempts.all())
self.assertEqual(len(voiceover["items"]), 2)
self.assertEqual([attempt.model_config_id for attempt in attempts], [primary.id, primary.id])
self.assertTrue(all(attempt.operation == "tts" for attempt in attempts))
self.assertEqual(self.direct_mock.synthesize.call_count, 2)
self.assertEqual(Asset.objects.filter(origin_task=task).count(), 2)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
def test_direct_model_with_outbound_disabled_retries_only_primary_and_releases(self):
primary = self.model(self.direct_provider, "voice-direct-locked", outbound=False)
candidate = self.model(
self.provider("voice-unused-candidate", 20),
"voice-unused",
outbound=False,
voice="alloy",
)
self.direct_mock.synthesize.side_effect = requests.ConnectionError("offline")
unused = Mock()
unused.synthesize.return_value = (b"unused", 1000)
self.provider_mocks[candidate.id] = unused
with self.assertRaises(requests.ConnectionError):
self.synthesize([{"index": 0, "text": "旁白"}])
task = AITask.objects.get(task_type=AITask.Type.VOICEOVER)
attempts = list(task.model_attempts.all())
self.assertEqual([attempt.model_config_id for attempt in attempts], [primary.id, primary.id])
self.assertFalse(unused.synthesize.called)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
def test_openai_candidate_enters_dynamically_and_uses_its_voice_map(self):
primary = self.model(self.direct_provider, "voice-dynamic-primary", outbound=True)
candidate = self.model(
self.provider("voice-dynamic-candidate", 20),
"voice-dynamic-fallback",
outbound=False,
voice="alloy",
base_cost="0.25",
)
self.direct_mock.synthesize.side_effect = requests.ConnectionError("offline")
fallback = Mock()
fallback.synthesize.return_value = (b"fallback-mp3", 1500)
self.provider_mocks[candidate.id] = fallback
voiceover = self.synthesize([{"index": 0, "text": "候选配音"}])
task = AITask.objects.get(task_type=AITask.Type.VOICEOVER)
attempts = list(task.model_attempts.all())
self.assertEqual(len(voiceover["items"]), 1)
self.assertEqual(
[attempt.model_config_id for attempt in attempts],
[primary.id, primary.id, candidate.id],
)
self.assertTrue(attempts[-1].is_fallback)
self.assertEqual(fallback.synthesize.call_args.kwargs["voice_type"], "alloy")
self.assertEqual(fallback.synthesize.call_args.kwargs["model"], candidate.name)
self.assertEqual(fallback.synthesize.call_args.kwargs["endpoint"], "audio/speech")
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
def test_all_candidates_fail_release_once_and_create_no_audio_assets(self):
primary = self.model(self.direct_provider, "voice-all-primary", outbound=True)
candidate = self.model(
self.provider("voice-all-candidate", 20),
"voice-all-fallback",
outbound=False,
voice="alloy",
)
self.direct_mock.synthesize.side_effect = requests.ConnectionError("offline")
failed = Mock()
failed.synthesize.side_effect = requests.ConnectionError("also offline")
self.provider_mocks[candidate.id] = failed
with self.assertRaises(requests.ConnectionError):
self.synthesize([{"index": 0, "text": "失败旁白"}])
task = AITask.objects.get(task_type=AITask.Type.VOICEOVER)
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(Asset.objects.filter(origin_task=task).count(), 0)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
class OpenAICompatibleSpeechTests(SimpleTestCase):
@patch("apps.ai.providers.openai_compatible.requests.post")
def test_standard_audio_speech_request(self, post):
response = Mock()
response.ok = True
response.content = b"mp3-bytes"
response.headers = {"x-audio-duration-ms": "1234"}
post.return_value = response
provider = OpenAICompatibleProvider(base_url="https://voice.example/v1", api_key="secret")
audio, duration_ms = provider.synthesize(
model="tts-model",
text="你好",
voice_type="alloy",
speed_ratio=1.2,
endpoint="audio/speech",
timeout=9,
)
self.assertEqual(audio, b"mp3-bytes")
self.assertEqual(duration_ms, 1234)
call = post.call_args
self.assertEqual(call.args[0], "https://voice.example/v1/audio/speech")
self.assertEqual(call.kwargs["json"]["voice"], "alloy")
self.assertEqual(call.kwargs["json"]["response_format"], "mp3")
self.assertEqual(call.kwargs["timeout"], 9)
+5 -5
View File
@@ -902,9 +902,9 @@ class StandaloneImageReferenceTests(TestCase):
self.assertIn("用户要求是最高优先级", used_prompt)
prov.image_generation.assert_not_called()
def test_refs_force_image_edit_model_when_volcano_selected(self):
"""用户选了火山 Seedream(只有图生图、不会真锁主体)却传了参考图时,必须自动切到支持
image_edit gpt-image-2否则参考图形同虚设回归保护选火山+传参考图=完全不参考"""
def test_refs_keep_user_selected_volcano_model_as_primary(self):
"""图片创作上传参考图后仍必须尊重用户选择的 Seedream;它可走带 image 的图生图,
只有真实调用失败后才允许统一路由按配置 Fallback不能在建任务阶段静默换成 gpt-image"""
from apps.ai.models import ModelProvider
vp, _ = ModelProvider.objects.get_or_create(name="volcengine", defaults={"display_name": "火山"})
@@ -919,8 +919,8 @@ class StandaloneImageReferenceTests(TestCase):
team=self.team, user=self.user, prompt="不同场景", mode="image", count=1,
image_model="volcano", reference_image_ids=[str(ref.id)],
)
# 任务最终落在支持 image_edit 的模型(gpt-image),而不是用户选的火山 Seedream
self.assertIn("gpt-image", tasks[0].model_config.name)
self.assertEqual(tasks[0].model_config.provider_id, vp.id)
self.assertIn("seedream", tasks[0].model_config.name)
def test_model_tryon_keeps_each_user_selected_model_even_with_extra_refs(self):
"""模特上身图由专用 Worker 向两类模型传商品/人物参考图,不能沿用自由创作的自动换模逻辑。"""
@@ -16,6 +16,13 @@ class Command(BaseCommand):
"status": ModelProvider.Status.ACTIVE,
},
)
# 只合并路由段,避免覆盖后台已维护的其他供应商 metadata。
provider_metadata = dict(provider.metadata or {})
provider_routing = dict(provider_metadata.get("routing") or {})
provider_routing.update(VOLCANO_PROVIDER["metadata"]["routing"])
provider_metadata["routing"] = provider_routing
provider.metadata = provider_metadata
provider.save(update_fields=["metadata"])
count = 0
for item in VOLCANO_MODELS:
@@ -16,6 +16,13 @@ class Command(BaseCommand):
"status": ModelProvider.Status.ACTIVE,
},
)
# 保留 api_version 等既有配置,仅同步目录声明的路由优先级。
provider_metadata = dict(provider.metadata or {})
provider_routing = dict(provider_metadata.get("routing") or {})
provider_routing.update(YUNQI_PROVIDER["metadata"]["routing"])
provider_metadata["routing"] = provider_routing
provider.metadata = provider_metadata
provider.save(update_fields=["metadata"])
count = 0
for item in YUNQI_MODELS:
+55
View File
@@ -258,6 +258,61 @@
overflow: auto;
}
.admin-json-err { color: var(--accent-crimson); background: var(--crimson-bg); border-color: var(--crimson-bd); }
.admin-attempt-list {
margin: 0 0 20px;
background: var(--surface);
border: 1px solid var(--border-muted);
border-radius: var(--r-md);
overflow: hidden;
}
.admin-attempt {
padding: 14px 16px;
}
.admin-attempt + .admin-attempt { border-top: 1px solid var(--border-faint); }
.admin-attempt-head {
display: flex;
align-items: center;
gap: 8px;
margin-bottom: 8px;
}
.admin-attempt-seq { color: var(--black-alpha-48); font-size: 11px; }
.admin-attempt-kind { margin-left: auto; color: var(--black-alpha-56); font-size: 11px; }
.admin-attempt-model {
display: flex;
align-items: baseline;
gap: 6px;
color: var(--accent-black);
font-size: 13.5px;
line-height: 1.5;
}
.admin-attempt-model strong { font-weight: 500; }
.admin-attempt-sep { color: var(--black-alpha-24); }
.admin-attempt-meta {
display: flex;
flex-wrap: wrap;
gap: 4px 14px;
margin-top: 6px;
color: var(--black-alpha-48);
font-size: 11px;
}
.admin-attempt-error {
display: flex;
flex-direction: column;
gap: 4px;
margin-top: 10px;
padding: 8px 10px;
background: var(--crimson-bg);
border-radius: var(--r-sm);
color: var(--accent-crimson);
font-size: 12px;
line-height: 1.6;
word-break: break-word;
}
@media (max-width: 640px) {
.admin-attempt-head { align-items: flex-start; flex-wrap: wrap; }
.admin-attempt-kind { width: 100%; margin-left: 0; }
}
/* ── 计费 / 额度策略 ── */
.admin-switch-row {
@@ -30,6 +30,12 @@ function statusPill(status: string) {
return <span className="pill info"><span className="dot" /></span>;
}
function attemptKind(attempt: AdminTaskDetail["attempts"][number]) {
if (attempt.is_fallback) return "Fallback";
if (attempt.is_retry) return "原模型重试";
return "首次调用";
}
export function AdminTasksPage({ notify }: { notify: Notify }) {
const [tasks, setTasks] = useState<AdminTask[]>([]);
const [count, setCount] = useState(0);
@@ -171,13 +177,55 @@ export function AdminTasksPage({ notify }: { notify: Notify }) {
<p className="admin-modal-desc"></p>
) : (
<>
{(() => {
const attempts = detail.attempts || [];
const finalAttempt = [...attempts].reverse().find((attempt) => attempt.status === "succeeded");
const fallbackUsed = attempts.some((attempt) => attempt.is_fallback);
return (
<div className="admin-detail-meta">
<div><span className="k"></span><span className="v">{statusPill(detail.status)}</span></div>
<div><span className="k"></span><span className="v">{detail.team_name || "—"}</span></div>
<div><span className="k"></span><span className="v">{detail.model_name || "—"}</span></div>
<div><span className="k"></span><span className="v mono">{pts(detail.estimated_cost)} {pts(detail.actual_cost)} {detail.cost_anomaly ? " ⚠" : ""}</span></div>
<div><span className="k"></span><span className="v mono">{Number(detail.base_cost || 0) > 0 ? `¥${detail.base_cost} · 毛利 ${detail.margin_yuan != null ? `¥${detail.margin_yuan}` : "—"}` : "未知"}</span></div>
<div><span className="k"></span><span className="v mono">{attempts.length ? `${attempts.length} 次 · ${fallbackUsed ? "发生 Fallback" : "未切换"}` : "旧任务 · 无尝试链"}</span></div>
{finalAttempt && <div><span className="k"></span><span className="v">{finalAttempt.provider_display_name || finalAttempt.provider_name} / {finalAttempt.model_display_name || finalAttempt.model_name}</span></div>}
</div>
);
})()}
{detail.attempts?.length > 0 && (
<>
<div className="admin-detail-subhead"> · {detail.attempts.length} </div>
<div className="admin-attempt-list">
{detail.attempts.map((attempt) => (
<div className="admin-attempt" key={attempt.id}>
<div className="admin-attempt-head">
<span className="admin-attempt-seq mono">// {String(attempt.sequence).padStart(2, "0")}</span>
{statusPill(attempt.status)}
<span className="admin-attempt-kind mono">{attemptKind(attempt)}</span>
</div>
<div className="admin-attempt-model">
<span>{attempt.provider_display_name || attempt.provider_name}</span>
<span className="admin-attempt-sep">/</span>
<strong>{attempt.model_display_name || attempt.model_name}</strong>
</div>
<div className="admin-attempt-meta mono">
<span>{attempt.operation}</span>
<span>{attempt.duration_ms == null ? "耗时 —" : `耗时 ${attempt.duration_ms} ms`}</span>
<span>{Number(attempt.platform_cost || 0) > 0 ? `平台成本 ¥${attempt.platform_cost}` : "平台成本未知"}</span>
{attempt.provider_task_id && <span>Provider ID {attempt.provider_task_id}</span>}
</div>
{(attempt.error_type || attempt.raw_error) && (
<div className="admin-attempt-error">
<span className="mono">[{attempt.error_type || "unknown"}{attempt.provider_error_code ? ` · ${attempt.provider_error_code}` : ""}]</span>
<span>{attempt.raw_error || attempt.safe_error_summary}</span>
</div>
)}
</div>
))}
</div>
</>
)}
{detail.error_message && (
<>
<div className="admin-detail-subhead"></div>
+28
View File
@@ -133,6 +133,33 @@ export type AdminTask = {
reapable?: boolean;
created_at: string;
};
export type AdminModelAttempt = {
id: string;
sequence: number;
provider_name: string;
provider_display_name: string;
model_name: string;
model_display_name: string;
public_model_name: string;
capability: string;
operation: string;
status: string;
is_retry: boolean;
is_fallback: boolean;
previous_attempt: string | null;
provider_task_id: string;
started_at: string;
finished_at: string | null;
duration_ms: number | null;
error_type: string;
provider_error_code: string;
raw_error: string;
safe_error_summary: string;
usage: Record<string, unknown>;
platform_cost: string;
request_summary: Record<string, unknown>;
response_summary: Record<string, unknown>;
};
export type AdminTaskDetail = AdminTask & {
project: string | null;
idempotency_key: string;
@@ -141,6 +168,7 @@ export type AdminTaskDetail = AdminTask & {
error_message: string;
submitted_at: string | null;
completed_at: string | null;
attempts: AdminModelAttempt[];
};
export type AdminLedger = {
id: string;
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,184 @@
# 模型调用与动态 Fallback 完成说明
> 状态:已完成
>
> 最后核对:2026-07-21
>
> 用途:说明当前模型调用、重试、动态切换、计费和日志流程。
## 1. 当前调用流程
```text
用户选择模型 / 系统默认模型
→ 创建一个 AITask
→ 预留一次积分
→ 调用主模型
→ 失败后按配置重试主模型
→ 仍失败且允许 Fallback
→ 实时查询已启用的兼容模型
→ 按动态顺序逐个尝试
→ 成功:保存结果并结算一次
→ 全部失败:释放一次预留
```
主模型始终是用户本次选择的模型;Fallback 不修改用户选择、默认模型或下次调用。
## 2. Retry 与 Fallback
- `Retry`:再次调用当前模型。
- `Fallback`:切换到其他兼容模型。
- `调用审计`:记录每次真实调用,不参与模型评分或自动学习。
当前默认值:
| 能力 | 主模型重试 | 单次超时 | 总时限 |
|---|---:|---:|---:|
| 文本 | 2 次,等待 1 / 3 秒 | 120 秒;流式 300 秒 | 480 秒 |
| 图片 | 1 次,等待 3 秒 | 300 秒 | 900 秒 |
| 配音 | 1 次,等待 2 秒 | 60 秒 | 180 秒 |
| 视频提交 | 1 次,等待 3 秒 | 120 秒 | 300 秒 |
| 视频成片 | 不重新提交 | 单次轮询 60 秒 | 1800 秒 |
单个逻辑任务最多尝试 3 个模型、发出 5 次真实请求。
## 3. 配置位置
全局策略在:
```text
core/backend/.env
core/backend/airshelf/settings/base.py → MODEL_ROUTING_POLICY
```
控制重试、超时、总时限、最多模型数、最多调用数和后处理退避。修改后必须同时重启 API 与 Celery Worker。
每个模型的开关保存在数据库 `ModelConfig.metadata.routing`
```json
{
"fallback_on_failure": false,
"fallback_candidate": true
}
```
- `fallback_on_failure`:当前模型失败后是否允许切到其他模型。
- `fallback_candidate`:当前模型是否允许被其他失败模型选中。
- 开关按模型独立配置,不在代码中写死供应商或模型名单。
## 4. 火山 / 豆包直连现状
文本、配音、视频都已接入统一 Fallback 代码,但当前直连模型统一配置为:
```text
fallback_on_failure = false
fallback_candidate = true
```
当前行为:
- 允许原模型重试。
- 不允许从直连模型主动切出。
- 已启用且能力匹配时,可以作为其他模型的候选。
- 以后只改模型 metadata 即可开启,无需重新开发路由代码。
## 5. 动态候选规则
候选每次失败后实时读取数据库,必须同时满足:
- 供应商已启用。
- 模型已启用。
- `fallback_candidate=true`
- 能力一致:`text``image``audio``video`
- 操作要求一致:流式、结构化、图片编辑、参考素材数量、比例、分辨率、时长、音色映射等。
- 本任务尚未尝试过。
排序规则:
```text
火山 / 豆包直连
→ YunQi
→ 其他 OpenAI API 兼容供应商
→ 同层级按 ModelConfig.updated_at 从新到旧
```
除现有火山 / 豆包专用适配器外,新增供应商只支持 OpenAI API 兼容协议。
## 6. 任务与计费
一次用户操作始终只有:
- 一个 `AITask`
- 一次积分预留。
- 一次最终扣费或释放。
每次真实模型调用只新增一条 `AIModelAttempt`,不会单独预留或扣费。
- Fallback 成功:按原逻辑任务结算一次。
- 全部失败:释放原预留,用户扣费为 0。
- 多次尝试产生的上游费用累计为平台成本,不转换成多笔用户扣费。
## 7. 调用日志
`AIModelAttempt` 记录:
- 调用顺序。
- 供应商与真实模型快照。
- 首次调用、重试或 Fallback。
- 成功 / 失败、耗时和错误分类。
- Provider 任务 ID、用量和平台成本。
- 脱敏请求 / 响应摘要。
管理员在 `/admin/tasks` 的任务详情查看完整调用链。普通用户端保持静默,不显示内部供应商、真实候选、失败原因或尝试次数。
公开别名保持不变:
- Gemini 3.1 Pro → `AirShelf Script`
- `gpt-image-2``AirShelf Image`
- 火山 / 豆包直连继续显示原公开名称。
## 8. 视频特殊规则
- 只有尚未取得远端任务 ID,且明确提交失败时,才允许重试或 Fallback。
- 提交读取超时或响应缺少任务 ID 视为“远端状态未知”,禁止重复提交。
- 获得远端任务 ID 后固定使用真正提交成功的模型轮询和结算。
- 普通轮询超时只继续轮询,不切模型、不重新提交。
- 达到成片总时限后才进入最终失败。
- 下载、上传、TOS 或资产保存失败只重试后处理,不重新调用模型。
## 9. 已接入入口
- 脚本生成、脚本优化、单镜优化。
- 显式实体提取、脚本实体同步兼容入口。
- 独立图片生成 / 编辑、模特上身图、平台套图。
- 商品 / 人物 / 场景基础资产、项目与模特三视图。
- 故事板分镜生成。
- 配音生成。
- 视频片段生成。
- 自由创作 / 免费视频生成。
## 10. 发布要求
- 数据库迁移:`0028_seed_model_routing_metadata``0029_aimodelattempt`
- API 与 Celery Worker 必须使用同一版后端镜像。
- API 与 Worker 必须读取同一份路由配置。
- 发布后至少冒烟验证文本、图片、配音、视频和 `/admin/tasks` 调用链。
- 回滚代码时不要直接反向执行 `0028`,避免删除已维护的模型路由 metadata。
## 11. 关键代码
```text
core/backend/apps/ai/routing_policy.py 全局策略读取与校验
core/backend/apps/ai/model_routing.py 能力匹配与动态候选排序
core/backend/apps/ai/routing_executor.py 重试、Fallback 与调用审计
core/backend/apps/ai/models.py AITask / AIModelAttempt
core/backend/apps/ai/services.py 各业务入口编排
core/backend/apps/ai/free_video.py 自由视频提交与轮询
core/backend/apps/ai/script_agent.py 脚本流式调用
```
## 12. 非本次范围
- 未修改 `/admin/providers` 页面。
- 未实现用户反馈学习、模型评分或自动调权。
- 未写死备用模型链。
- 未支持新增非 OpenAI API 私有协议供应商。