diff --git a/core/backend/.env b/core/backend/.env index d77d76a..e70a40f 100644 --- a/core/backend/.env +++ b/core/backend/.env @@ -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 diff --git a/core/backend/.env.example b/core/backend/.env.example index cd12f60..ca22a2a 100644 --- a/core/backend/.env.example +++ b/core/backend/.env.example @@ -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 diff --git a/core/backend/ARCHITECTURE.md b/core/backend/ARCHITECTURE.md index 2b9e33d..a70cf0f 100644 --- a/core/backend/ARCHITECTURE.md +++ b/core/backend/ARCHITECTURE.md @@ -1,8 +1,8 @@ # AirShelf 后端架构 > 适用目录:`core/backend`。 -> 最后核对:2026-07-18。 - +> 最后核对:2026-07-21。 + > 本文描述当前代码结构;系统级拓扑见 [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 diff --git a/core/backend/airshelf/settings/base.py b/core/backend/airshelf/settings/base.py index 0d8968d..9437e6b 100644 --- a/core/backend/airshelf/settings/base.py +++ b/core/backend/airshelf/settings/base.py @@ -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。 + # 合法范围 1~20,且不得小于 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) diff --git a/core/backend/apps/adminpanel/serializers.py b/core/backend/apps/adminpanel/serializers.py index 8af33a6..a37a366 100644 --- a/core/backend/apps/adminpanel/serializers.py +++ b/core/backend/apps/adminpanel/serializers.py @@ -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) diff --git a/core/backend/apps/adminpanel/tests.py b/core/backend/apps/adminpanel/tests.py index a7282c5..4ed8dd4 100644 --- a/core/backend/apps/adminpanel/tests.py +++ b/core/backend/apps/adminpanel/tests.py @@ -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: diff --git a/core/backend/apps/adminpanel/views.py b/core/backend/apps/adminpanel/views.py index d969669..57b18d0 100644 --- a/core/backend/apps/adminpanel/views.py +++ b/core/backend/apps/adminpanel/views.py @@ -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) diff --git a/core/backend/apps/ai/apps.py b/core/backend/apps/ai/apps.py index a89e3fc..48d843b 100644 --- a/core/backend/apps/ai/apps.py +++ b/core/backend/apps/ai/apps.py @@ -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 + diff --git a/core/backend/apps/ai/catalog.py b/core/backend/apps/ai/catalog.py index b6b4dc7..6ba0021 100644 --- a/core/backend/apps/ai/catalog.py +++ b/core/backend/apps/ai/catalog.py @@ -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) + diff --git a/core/backend/apps/ai/free_video.py b/core/backend/apps/ai/free_video.py index da93ee3..65fd871 100644 --- a/core/backend/apps/ai/free_video.py +++ b/core/backend/apps/ai/free_video.py @@ -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 diff --git a/core/backend/apps/ai/migrations/0028_seed_model_routing_metadata.py b/core/backend/apps/ai/migrations/0028_seed_model_routing_metadata.py new file mode 100644 index 0000000..9efb30b --- /dev/null +++ b/core/backend/apps/ai/migrations/0028_seed_model_routing_metadata.py @@ -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)] diff --git a/core/backend/apps/ai/migrations/0029_aimodelattempt.py b/core/backend/apps/ai/migrations/0029_aimodelattempt.py new file mode 100644 index 0000000..7840e9b --- /dev/null +++ b/core/backend/apps/ai/migrations/0029_aimodelattempt.py @@ -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") + ], + }, + ), + ] diff --git a/core/backend/apps/ai/model_routing.py b/core/backend/apps/ai/model_routing.py new file mode 100644 index 0000000..8387140 --- /dev/null +++ b/core/backend/apps/ai/model_routing.py @@ -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", +] diff --git a/core/backend/apps/ai/models.py b/core/backend/apps/ai/models.py index f81d6cd..d4fa8e6 100644 --- a/core/backend/apps/ai/models.py +++ b/core/backend/apps/ai/models.py @@ -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 的写死值**(零回归)。 diff --git a/core/backend/apps/ai/providers/openai_compatible.py b/core/backend/apps/ai/providers/openai_compatible.py index a6fbed8..edb305e 100644 --- a/core/backend/apps/ai/providers/openai_compatible.py +++ b/core/backend/apps/ai/providers/openai_compatible.py @@ -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 diff --git a/core/backend/apps/ai/providers/volcano.py b/core/backend/apps/ai/providers/volcano.py index 037a341..f5847cd 100644 --- a/core/backend/apps/ai/providers/volcano.py +++ b/core/backend/apps/ai/providers/volcano.py @@ -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() diff --git a/core/backend/apps/ai/providers/yunqi.py b/core/backend/apps/ai/providers/yunqi.py index b650d44..76cbe7b 100644 --- a/core/backend/apps/ai/providers/yunqi.py +++ b/core/backend/apps/ai/providers/yunqi.py @@ -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() diff --git a/core/backend/apps/ai/routing_executor.py b/core/backend/apps/ai/routing_executor.py new file mode 100644 index 0000000..8cd7dbb --- /dev/null +++ b/core/backend/apps/ai/routing_executor.py @@ -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", +] diff --git a/core/backend/apps/ai/routing_policy.py b/core/backend/apps/ai/routing_policy.py new file mode 100644 index 0000000..36827e9 --- /dev/null +++ b/core/backend/apps/ai/routing_policy.py @@ -0,0 +1,303 @@ +"""模型路由策略的集中读取与校验。 + +本模块只负责把 Django settings 中的 ``MODEL_ROUTING_POLICY`` 转成不可变配置对象, +不执行重试、Fallback、Provider 调用或账务操作。业务入口后续只消费这里返回的策略, +不得各自维护 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", +] diff --git a/core/backend/apps/ai/script_agent.py b/core/backend/apps/ai/script_agent.py index 122048a..925d5e4 100644 --- a/core/backend/apps/ai/script_agent.py +++ b/core/backend/apps/ai/script_agent.py @@ -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)。计费 reserve→charge/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 diff --git a/core/backend/apps/ai/services.py b/core/backend/apps/ai/services.py index 525820f..798c4fb 100644 --- a/core/backend/apps/ai/services.py +++ b/core/backend/apps/ai/services.py @@ -19,12 +19,14 @@ from django.utils import timezone from apps.ai.models import AITask, ModelConfig from apps.ai.generation_errors import TASK_OPERATIONS, classify_generation_error, public_error_for_task +from apps.ai.model_routing import ModelRequirements, capability_metadata, model_allows_fallback from apps.ai.providers import ( OpenAICompatibleProvider, TtsNotConfigured, VolcanoArkProvider, VolcanoTtsProvider, ) +from apps.ai.routing_executor import AttemptMetadata, execute_model_call from apps.assets.models import Asset, AssetFile from apps.assets.storage import TosStorage from apps.billing.services.ledger import charge_reserved_credit, release_credit, reserve_credit @@ -80,7 +82,19 @@ def resolve_image_model(key: str | None) -> "ModelConfig | None": # 火山官方直连(SeeDream 生图 / Seedance 视频 / 豆包文本)走 ARK SDK;其余 provider 一律 # 视为「OpenAI 兼容中转站」走通用适配器。加/换中转站 = DB 加一行 ModelProvider,零改代码。 # 注意:DB 里火山 provider 实际命名为 "volcengine"(豆包),必须包含,否则会被错路由到中转站。 -OFFICIAL_DIRECT_PROVIDERS = {"volcengine", "volcano", "ark", "volcano_ark"} +OFFICIAL_DIRECT_PROVIDERS = {"volcengine", "volcano", "ark", "volcano_ark", "doubao"} + + +def public_model_name(model_config: ModelConfig) -> str: + """普通用户公开名称保持稳定;Fallback 的真实模型只在管理员尝试链中展示。""" + + if model_config.provider.name in OFFICIAL_DIRECT_PROVIDERS: + return model_config.display_name + if model_config.capability == ModelConfig.Capability.TEXT: + return "AirShelf Script" + if model_config.capability == ModelConfig.Capability.IMAGE: + return "AirShelf Image" + return model_config.display_name def resolve_provider_credentials(provider) -> tuple[str | None, str | None]: @@ -113,6 +127,114 @@ def get_image_provider(model_config: ModelConfig): return build_provider(model_config) +def execute_routed_image_request( + *, + task: AITask, + primary_model: ModelConfig, + prompt: str, + reference_images: list[str], + aspect_ratio: str | None = None, + edit_size: str | None = None, + direct_size: str | None = None, + generate_size: str | None = None, + request_summary: dict | None = None, +): + """图片入口共用的模型调用薄层:只处理能力声明、Provider 调用和尝试成本审计。 + + 业务入口仍负责提示词、参考图顺序、尺寸、任务/资产/账务终态;这里不创建任务、不结算积分, + 因而图片创作、上身图、套图和项目基础资产可以复用而不互相耦合。 + """ + from apps.billing.pricing import quote_flat + + references = list(reference_images or []) + reference_count = len(references) + reference_mode = "none" + if reference_count == 1: + reference_mode = "single" + elif reference_count > 1: + reference_mode = "multiple" + operation = "image_edit" if reference_count else "image_generate" + requirements = ModelRequirements( + capability=ModelConfig.Capability.IMAGE, + operation=operation, + reference_mode=reference_mode, + reference_images=reference_count, + aspect_ratio=aspect_ratio or None, + ) + + def invoke_image(actual_model: ModelConfig, timeout: float): + actual_provider = get_image_provider(actual_model) + if references and hasattr(actual_provider, "image_edit"): + kwargs = { + "model": actual_model.name, + "prompt": prompt, + "images": references, + "timeout": timeout, + } + if edit_size: + kwargs["size"] = edit_size + actual_response = actual_provider.image_edit(**kwargs) + elif references: + kwargs = { + "model": actual_model.name, + "endpoint": actual_model.endpoint, + "prompt": prompt, + "image": references, + "timeout": timeout, + } + if direct_size or edit_size: + kwargs["size"] = direct_size or edit_size + actual_response = actual_provider.image_generation(**kwargs) + else: + kwargs = { + "model": actual_model.name, + "endpoint": actual_model.endpoint, + "prompt": prompt, + "timeout": timeout, + } + if generate_size: + kwargs["size"] = generate_size + actual_response = actual_provider.image_generation(**kwargs) + try: + actual_media = actual_provider.extract_first_media_url(actual_response) + except Exception as exc: + # Provider 已返回并可能产生上游费用;响应解析失败仍需审计本次真实尝试成本。 + candidate_quote = quote_flat(actual_model, team=task.team) + exc.attempt_metadata = AttemptMetadata( + usage=actual_response.get("usage") if isinstance(actual_response, dict) else {}, + platform_cost=candidate_quote.base_cost_yuan, + response_summary={"media_missing": True}, + ) + raise + return actual_response, actual_media + + def image_result_metadata(result, actual_model: ModelConfig): + actual_response, _ = result + candidate_quote = quote_flat(actual_model, team=task.team) + return AttemptMetadata( + usage=actual_response.get("usage") if isinstance(actual_response, dict) else {}, + platform_cost=candidate_quote.base_cost_yuan, + response_summary={"media_found": True}, + ) + + summary = { + "operation": operation, + "prompt_length": len(prompt), + "aspect_ratio": aspect_ratio or None, + "reference_images": reference_count, + } + summary.update(request_summary or {}) + return execute_model_call( + task=task, + primary_model=primary_model, + requirements=requirements, + public_model_name=public_model_name(primary_model), + invoke=invoke_image, + request_summary=summary, + result_metadata=image_result_metadata, + ) + + def get_text_provider(model_config: ModelConfig): return build_provider(model_config) @@ -121,6 +243,13 @@ def get_video_provider(model_config: ModelConfig): return build_provider(model_config) +def get_audio_provider(model_config: ModelConfig): + """豆包/火山现有 TTS 走专用直连;其余启用模型统一走 OpenAI 兼容 ``audio/speech``。""" + if model_config.provider.name in OFFICIAL_DIRECT_PROVIDERS: + return VolcanoTtsProvider() + return build_provider(model_config) + + # estimate_cost() 已退役:全平台定价统一走 apps/billing/pricing.py 计价引擎(积分制)。 # flat 类型默认价由 create_ai_task 内 quote_flat 提供;视频/配音各入口自带 quote。 @@ -197,12 +326,25 @@ def _coerce_tag_entries(items: object, limit: int = 6) -> tuple[list[str], dict[ return names, prompts +def _parse_cast_scene_response(text: str) -> dict: + """解析旧版脚本后人物/场景提取结果,并把结构校验纳入单次路由尝试。""" + match = re.search(r"\{.*\}", text or "", re.DOTALL) # 容忍 markdown / 前后解释文字 + if not match: + raise ValueError("人物与场景提取结果不是有效 JSON") + data = json.loads(match.group(0)) + if not isinstance(data, dict): + raise ValueError("人物与场景提取结果必须是 JSON 对象") + cast, cast_prompts = _coerce_tag_entries(data.get("cast")) + scenes, scene_prompts = _coerce_tag_entries(data.get("scenes")) + return {"cast": cast, "scenes": scenes, "cast_prompts": cast_prompts, "scene_prompts": scene_prompts} + + def extract_cast_and_scenes(*, project, user, content: str) -> dict: """轻量调一次文本模型,从脚本里抽取人物 / 场景标签及每个标签的建议生图提示词。 - 出稿后的增益步骤(对齐流程文档「自动从脚本提取信息」),走完整的 AITask + 计费闭环 - (reserve→charge/release,任务类型记为 script_optimization)。但全程 best-effort: - 无可用模型 / 预扣失败 / 调用失败 / 解析失败都吞掉返回空,绝不阻断脚本生成主流程。 + 这是存量兼容入口;当前 Script Agent 已在同一次结构化出稿中携带 entities,不会额外调用 + 本函数。若旧调用方或未来流程重新启用它,仍走统一重试/Fallback、尝试日志和一次性计费 + 闭环。全程 best-effort:无可用模型、预扣失败、调用或解析失败都返回空,绝不阻断主流程。 返回 {cast, scenes, cast_prompts, scene_prompts}。 """ empty = {"cast": [], "scenes": [], "cast_prompts": {}, "scene_prompts": {}} @@ -217,22 +359,38 @@ def extract_cast_and_scenes(*, project, user, content: str) -> dict: user=user, task_type=AITask.Type.SCRIPT_OPTIMIZATION, model_config=model_config, - request_payload={"model": model_config.name, "endpoint": model_config.endpoint, "messages": messages}, + request_payload={ + "model": model_config.name, + "endpoint": model_config.endpoint, + "messages": messages, + "model_routing_v1": True, + }, ) except Exception: return empty # 余额不足等预扣失败:跳过提取,不挡出稿 reservation = task.credit_reservation + # 每条真实尝试的成本由统一执行器按实际模型累加;用户积分仍只按逻辑任务结算一次。 + 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) - text = provider.extract_text(response) + routed = execute_routed_text_request( + task=task, + primary_model=model_config, + messages=messages, + streaming=False, + structured_output=True, + business_operation="entity_extract", + temperature=0.3, + validate_text=_parse_cast_scene_response, + request_summary={"source": "legacy_post_script_extract"}, + ) + _text, response, parsed = routed.value - # LLM 调用已真实消耗 token → 按成功计费(无论后续能否解析出标签) with transaction.atomic(): task.status = AITask.Status.SUCCEEDED task.response_payload = response @@ -240,14 +398,7 @@ def extract_cast_and_scenes(*, project, user, content: str) -> dict: task.completed_at = timezone.now() task.save(update_fields=["status", "response_payload", "actual_cost", "completed_at", "updated_at"]) charge_reserved_credit(reservation=reservation, actual_amount=task.actual_cost) - - match = re.search(r"\{.*\}", text, re.DOTALL) # 容忍模型多裹了 markdown / 解释文字 - if not match: - return empty - data = json.loads(match.group(0)) - cast, cast_prompts = _coerce_tag_entries(data.get("cast")) - scenes, scene_prompts = _coerce_tag_entries(data.get("scenes")) - return {"cast": cast, "scenes": scenes, "cast_prompts": cast_prompts, "scene_prompts": scene_prompts} + return parsed except Exception as exc: with transaction.atomic(): task.status = AITask.Status.FAILED @@ -346,6 +497,22 @@ def _normalize_extracted_segment_refs(items: object, valid_ids: set[str]) -> lis return out +def _parse_extracted_entities_response(text: str) -> tuple[list[dict], list[dict]]: + """解析并校验实体提取结构;放进单次模型尝试内部,格式漂移可按统一策略重试/切换。""" + match = re.search(r"\{.*\}", text or "", re.DOTALL) + if not match: + raise ValueError("提取结果解析失败(模型未返回有效 JSON),请重试") + try: + data = json.loads(match.group(0)) + except json.JSONDecodeError as exc: + raise ValueError("提取结果解析失败,请重试") from exc + entities = _normalize_extracted_entities(data.get("entities")) + if not entities: + raise ValueError("没有从脚本里识别到角色 / 场景,可调整脚本后重试") + seg_refs = _normalize_extracted_segment_refs(data.get("segments"), {e["id"] for e in entities}) + return entities, seg_refs + + # 提取步固定锁定豆包 2.0 Pro(与脚本生成同款),不靠 get_default_model 的「最早创建」排序—— # 避免不同环境 DB 创建序漂移把提取路由到别的(中转站)推理模型。取不到再回落默认文本模型。 EXTRACT_TEXT_MODEL_NAME = "doubao-seed-2-0-pro-260215" @@ -367,7 +534,15 @@ def _resolve_extract_model_config(): return pinned or get_default_model(ModelConfig.Capability.TEXT) -def _collect_extract_text(provider, model_config, messages) -> tuple[str, dict]: +def _collect_extract_text( + provider, + model_config, + messages, + *, + temperature: float = 0.3, + timeout: float = 300, + on_event=None, +) -> tuple[str, dict]: """走与「脚本生成」同一条已在生产验证稳定的流式通道把模型输出收全。 豆包 seed-pro / GPT / Gemini 等思考模型,思考期只发 reasoning_content、正文期才发 content, @@ -381,8 +556,11 @@ def _collect_extract_text(provider, model_config, messages) -> tuple[str, dict]: model=model_config.name, endpoint=model_config.endpoint, messages=messages, - temperature=0.3, # 结构化抽取要稳:低温降低 JSON 漂移 + temperature=temperature, # 结构化抽取要稳:低温降低 JSON 漂移 + timeout=timeout, ): + if on_event is not None: + on_event(ev) etype = ev.get("type") if etype == "delta": content_parts.append(ev.get("text") or "") @@ -402,6 +580,407 @@ def _collect_extract_text(provider, model_config, messages) -> tuple[str, dict]: return text, payload +def execute_routed_text_request( + *, + task: AITask, + primary_model: ModelConfig, + messages: list[dict], + streaming: bool, + structured_output: bool, + business_operation: str, + temperature: float = 0.3, + validate_text=None, + request_summary: dict | None = None, + stream_event_callback=None, + abort_check=None, +): + """文本入口共用的模型调用薄层:能力声明、Provider 调用、输出校验和尝试成本审计。 + + ``validate_text`` 在一次真实调用内部执行;模型返回空文或无效结构时会进入同一重试/Fallback + 策略,而不是先把 Provider 调用判成功、随后在业务层直接失败。任务终态、积分和业务落库仍由 + 调用方负责。 + """ + from apps.billing.pricing import quote_flat + + features = set() + if streaming: + features.add("streaming") + if structured_output: + features.add("structured_output") + requirements = ModelRequirements( + capability=ModelConfig.Capability.TEXT, + operation="chat", + features=frozenset(features), + ) + stream_call_number = 0 + + def invoke_text(actual_model: ModelConfig, timeout: float): + nonlocal stream_call_number + stream_call_number += 1 + if abort_check is not None: + abort_check() + provider = get_text_provider(actual_model) + if streaming: + text, response = _collect_extract_text( + provider, + actual_model, + messages, + temperature=temperature, + timeout=timeout, + on_event=( + (lambda event: stream_event_callback(event, stream_call_number)) + if stream_event_callback is not None + else None + ), + ) + else: + response = provider.chat_completion( + model=actual_model.name, + endpoint=actual_model.endpoint, + messages=messages, + timeout=timeout, + ) + text = provider.extract_text(response) + if abort_check is not None: + abort_check() + try: + validated = validate_text(text) if validate_text is not None else None + except Exception as exc: + # Provider 已完成生成并可能产生上游费用;结构校验失败也必须把该次真实成本记入尝试。 + candidate_quote = quote_flat(actual_model, team=task.team) + exc.attempt_metadata = AttemptMetadata( + usage=response.get("usage") if isinstance(response, dict) else {}, + platform_cost=candidate_quote.base_cost_yuan, + response_summary={ + "streamed": streaming, + "structured_output": structured_output, + "content_chars": len(text or ""), + "validation_failed": True, + }, + ) + raise + return text, response, validated + + def text_result_metadata(result, actual_model: ModelConfig): + text, response, _ = result + candidate_quote = quote_flat(actual_model, team=task.team) + return AttemptMetadata( + usage=response.get("usage") if isinstance(response, dict) else {}, + platform_cost=candidate_quote.base_cost_yuan, + response_summary={ + "streamed": streaming, + "structured_output": structured_output, + "content_chars": len(text or ""), + }, + ) + + summary = { + "operation": "chat", + "business_operation": business_operation, + "message_count": len(messages), + "input_chars": sum(len(str(message.get("content") or "")) for message in messages), + "streaming": streaming, + "structured_output": structured_output, + } + summary.update(request_summary or {}) + return execute_model_call( + task=task, + primary_model=primary_model, + requirements=requirements, + public_model_name=public_model_name(primary_model), + invoke=invoke_text, + request_summary=summary, + result_metadata=text_result_metadata, + error_classifier=lambda exc, model: classify_generation_error( + exc, + operation=business_operation, + provider_name=model.provider.name, + internal_kind="processing_failed" if isinstance(exc, _RoutedTextStreamCancelled) else "", + reference_id=str(task.id), + ), + ) + + +class _RoutedTextStreamCancelled(RuntimeError): + """HTTP 客户端已断开;终止后台流读取,不再继续重试或切换模型。""" + + +def stream_routed_text_request(**kwargs): + """把统一文本执行器转换为可实时转发 Provider 事件的生成器。 + + 路由与尝试日志仍完全复用 ``execute_routed_text_request``。执行器运行在一个短生命周期 + 后台线程中,当前生成器从队列逐个转发首轮事件;若首轮失败后重试/Fallback,为避免普通 + 用户看到重复半截文本,后续轮次静默收全,只返回最终结构化结果。生成器返回值是 + ``RoutingExecutionResult``,调用方可用 ``yield from`` 或捕获 ``StopIteration.value`` 获取。 + """ + from queue import SimpleQueue + from threading import Event, Thread + + from django.db import close_old_connections + + queue = SimpleQueue() + cancelled = Event() + + def abort_check(): + if cancelled.is_set(): + raise _RoutedTextStreamCancelled("stream aborted (client disconnected)") + + def forward_event(event, call_number): + abort_check() + if call_number == 1: + queue.put(("event", event)) + + def worker(): + close_old_connections() + try: + result = execute_routed_text_request( + **kwargs, + stream_event_callback=forward_event, + abort_check=abort_check, + ) + except BaseException as exc: # noqa: BLE001 — 跨线程原样交回请求生成器处理 + queue.put(("error", exc)) + else: + queue.put(("done", result)) + finally: + close_old_connections() + + thread = Thread(target=worker, name=f"ai-text-stream-{kwargs['task'].id}", daemon=True) + thread.start() + try: + while True: + kind, value = queue.get() + if kind == "event": + yield value + elif kind == "error": + raise value + else: + return value + finally: + cancelled.set() + + +def execute_routed_audio_request( + *, + task: AITask, + primary_model: ModelConfig, + text: str, + public_voice: str, + speed_ratio: float, + user_id: str, + request_summary: dict | None = None, +): + """配音单句调用薄层:动态音色映射、OpenAI 兼容候选、尝试日志和实际成本。""" + from apps.billing.pricing import quote_voiceover + + char_count = len(text) + requirements = ModelRequirements( + capability=ModelConfig.Capability.AUDIO, + operation="tts", + language="zh-CN", + public_voice=public_voice, + char_count=char_count, + speed_ratio=float(speed_ratio or 1.0), + output_format="mp3", + ) + + def invoke_audio(actual_model: ModelConfig, timeout: float): + provider = get_audio_provider(actual_model) + if hasattr(provider, "configured") and not provider.configured: + raise TtsNotConfigured("语音合成供应商凭证未配置") + voice_map = capability_metadata(actual_model).get("voice_map") + actual_voice = ( + voice_map.get(public_voice) + if isinstance(voice_map, dict) and voice_map.get(public_voice) + else public_voice + ) + kwargs = { + "text": text, + "voice_type": actual_voice, + "speed_ratio": speed_ratio, + "uid": user_id, + "timeout": timeout, + } + if actual_model.provider.name not in OFFICIAL_DIRECT_PROVIDERS: + kwargs.update( + { + "model": actual_model.name, + "endpoint": actual_model.endpoint or "audio/speech", + "output_format": "mp3", + } + ) + audio, duration_ms = provider.synthesize(**kwargs) + if not isinstance(audio, (bytes, bytearray)) or not audio: + raise ValueError("语音合成未返回有效音频") + return bytes(audio), max(0, int(duration_ms or 0)) + + def audio_result_metadata(result, actual_model: ModelConfig): + audio, duration_ms = result + candidate_quote = quote_voiceover(actual_model, char_count=char_count, team=task.team) + return AttemptMetadata( + usage={"characters": char_count}, + platform_cost=candidate_quote.base_cost_yuan, + response_summary={ + "audio_bytes": len(audio), + "duration_ms": duration_ms, + "output_format": "mp3", + }, + ) + + summary = { + "operation": "tts", + "public_voice": public_voice, + "char_count": char_count, + "speed_ratio": float(speed_ratio or 1.0), + "language": "zh-CN", + "output_format": "mp3", + } + summary.update(request_summary or {}) + return execute_model_call( + task=task, + primary_model=primary_model, + requirements=requirements, + public_model_name=public_model_name(primary_model), + invoke=invoke_audio, + request_summary=summary, + result_metadata=audio_result_metadata, + error_classifier=lambda exc, model: classify_generation_error( + exc, + operation="voiceover_generate", + provider_name=model.provider.name, + reference_id=str(task.id), + ), + ) + + +class VideoSubmissionStateUnknown(RuntimeError): + """视频提交可能已到达供应商但未拿到可靠任务 ID;禁止自动重提。""" + + +def execute_routed_video_submit( + *, + task: AITask, + primary_model: ModelConfig, + prompt: str, + duration: int, + ratio: str, + resolution: str, + reference_images: list[str] | None = None, + content_items: list[dict] | None = None, + pricing_references: list[dict] | None = None, + generate_audio: bool = True, + seed: int | None = None, + search_mode: str = "off", + request_summary: dict | None = None, +): + """异步视频只路由“提交”阶段;拿到 Provider 任务 ID 后固定该实际模型轮询。""" + from apps.billing.pricing import quote_video_estimate + + references = list(reference_images or []) + routed_content_items = list(content_items) if content_items is not None else None + pricing_refs = list(pricing_references or []) + item_types = [str((item or {}).get("type") or "") for item in routed_content_items or []] + image_count = len(references) + item_types.count("image_url") + video_count = item_types.count("video_url") + audio_count = item_types.count("audio_url") + requirements = ModelRequirements( + capability=ModelConfig.Capability.VIDEO, + operation="video_generate", + features=frozenset({"generate_audio"}) if generate_audio else frozenset(), + reference_images=image_count, + reference_videos=video_count, + reference_audios=audio_count, + aspect_ratio=ratio, + resolution=resolution, + duration=duration, + ) + + def candidate_quote(actual_model: ModelConfig): + _tokens, quote = quote_video_estimate( + actual_model, + aspect_ratio=ratio, + resolution=resolution, + duration=duration, + references=pricing_refs, + team=task.team, + ) + return quote + + def invoke_video(actual_model: ModelConfig, timeout: float): + provider = get_video_provider(actual_model) + try: + response = provider.create_video_task( + model=actual_model.name, + endpoint=actual_model.endpoint, + prompt=prompt, + duration=duration, + ratio=ratio, + resolution=resolution, + reference_images=references or None, + generate_audio=generate_audio, + content_items=routed_content_items, + seed=seed, + search_mode=search_mode, + timeout=timeout, + ) + except requests.ReadTimeout as exc: + # 请求可能已被供应商接收;盲目重提会生成两条视频任务并产生双份上游成本。 + raise VideoSubmissionStateUnknown("视频提交响应超时,远端创建状态未知") from exc + provider_task_id = str(response.get("id") or response.get("task_id") or "") + if not provider_task_id: + exc = VideoSubmissionStateUnknown("视频提交响应缺少任务 ID,远端创建状态未知") + quote = candidate_quote(actual_model) + exc.attempt_metadata = AttemptMetadata( + usage=response.get("usage") if isinstance(response, dict) else {}, + platform_cost=quote.base_cost_yuan, + response_summary={"provider_task_id_missing": True}, + ) + raise exc + return response, provider_task_id + + def video_result_metadata(result, actual_model: ModelConfig): + response, provider_task_id = result + quote = candidate_quote(actual_model) + return AttemptMetadata( + provider_task_id=provider_task_id, + usage=response.get("usage") if isinstance(response, dict) else {}, + platform_cost=quote.base_cost_yuan, + response_summary={ + "remote_status": str(response.get("status") or "") if isinstance(response, dict) else "", + "provider_task_id_received": True, + }, + ) + + summary = { + "operation": "video_generate", + "prompt_length": len(prompt), + "duration": duration, + "aspect_ratio": ratio, + "resolution": resolution, + "reference_images": image_count, + "reference_videos": video_count, + "reference_audios": audio_count, + "generate_audio": generate_audio, + } + summary.update(request_summary or {}) + return execute_model_call( + task=task, + primary_model=primary_model, + requirements=requirements, + public_model_name=public_model_name(primary_model), + invoke=invoke_video, + request_summary=summary, + result_metadata=video_result_metadata, + error_classifier=lambda exc, model: classify_generation_error( + exc, + operation="video_generate", + provider_name=model.provider.name, + internal_kind="processing_failed" if isinstance(exc, VideoSubmissionStateUnknown) else "", + reference_id=str(task.id), + ), + ) + + # 在途状态:据此判「已有提取在跑」(提交侧防重复扣费 + 前端刷新后重建 loading) _EXTRACT_INFLIGHT = ( AITask.Status.CREATED, @@ -502,11 +1081,15 @@ def submit_extract_entities(*, project, user) -> AITask: "endpoint": model_config.endpoint, "messages": messages, "script_id": str(script.id), + "model_routing_v1": True, }, ) except Exception as exc: # 余额不足等预扣失败 raise ValueError("额度不足,无法提取(请先充值)") from exc + # 真实平台成本由每条 AIModelAttempt 按实际模型累加,避免 Fallback 后仍记主模型旧成本。 + task.base_cost = Decimal("0") + task.save(update_fields=["base_cost", "updated_at"]) extract_entities_task.delay(str(task.id)) return task @@ -533,20 +1116,25 @@ def run_extract_entities_task(*, task_id: str) -> None: task.submitted_at = timezone.now() task.save(update_fields=["status", "submitted_at", "updated_at"]) - provider = build_provider(model_config) - text, response = _collect_extract_text(provider, model_config, messages) - - match = re.search(r"\{.*\}", text, re.DOTALL) - if not match: - raise ValueError("提取结果解析失败(模型未返回有效 JSON),请重试") - try: - data = json.loads(match.group(0)) - except json.JSONDecodeError as exc: # 抓出来给用户可读话术,别把裸解析报错塞进 error_message - raise ValueError("提取结果解析失败,请重试") from exc - entities = _normalize_extracted_entities(data.get("entities")) - if not entities: - raise ValueError("没有从脚本里识别到角色 / 场景,可调整脚本后重试") - seg_refs = _normalize_extracted_segment_refs(data.get("segments"), {e["id"] for e in entities}) + if payload.get("model_routing_v1"): + routed = execute_routed_text_request( + task=task, + primary_model=model_config, + messages=messages, + streaming=True, + structured_output=True, + business_operation="entity_extract", + temperature=0.3, + validate_text=_parse_extracted_entities_response, + request_summary={"script_id": str(payload.get("script_id") or "")}, + ) + _text, response, parsed = routed.value + entities, seg_refs = parsed + else: + # 存量未迁移任务继续沿用旧调用,避免部署切换期间改变已排队任务行为。 + provider = build_provider(model_config) + text, response = _collect_extract_text(provider, model_config, messages) + entities, seg_refs = _parse_extracted_entities_response(text) script_id = payload.get("script_id") script = ScriptVersion.objects.filter(id=script_id).first() if script_id else None @@ -1391,7 +1979,6 @@ def generate_base_asset(*, project, user, kind: str, prompt: str, label: str = " model_config = get_default_model(ModelConfig.Capability.IMAGE) if model_config is None: raise ValueError("no active image model configured") - provider = get_image_provider(model_config) # 商品三视图:有真实商品主图 → 走 image_edit 以主图为参考,锁定包装(品牌字/配色/外形/Logo)一致; # 无主图或当前模型不支持 image_edit → 回落纯文生图(仅凭商品名脑补,不保证还原真实包装)。 product_ref_url = _product_cover_url(project.product) if kind == BaseAssetGroup.Kind.PRODUCT else "" @@ -1403,7 +1990,9 @@ def generate_base_asset(*, project, user, kind: str, prompt: str, label: str = " if ref_asset is not None: person_ref_url = _asset_preview_url(ref_asset) ref_url = product_ref_url or person_ref_url - use_edit = bool(ref_url) and hasattr(provider, "image_edit") + # 是否需要参考图由业务素材决定,不再用主 Provider 是否暴露 image_edit 来预判。 + # 直连 Seedream 可通过 image_generation(image=...) 使用同一参考图;候选资格由统一能力契约过滤。 + use_edit = bool(ref_url) if use_edit and product_ref_url: gen_prompt = build_product_triview_prompt_refs(project.product, prompt) elif use_edit and person_ref_url: @@ -1424,7 +2013,7 @@ def generate_base_asset(*, project, user, kind: str, prompt: str, label: str = " payload = { "model": model_config.name, "endpoint": model_config.endpoint, "prompt": gen_prompt, "kind": kind, "label": label or "", "group_id": str(group_id) if group_id else "", - "use_edit": use_edit, "reference_image": ref_url, + "use_edit": use_edit, "reference_image": ref_url, "model_routing_v1": True, # 角色立绘不再自动接力三视图;三视图只通过角色详情里的显式按钮生成。 # 保留字段为 False,兼容旧前端/旧任务读取,但不允许再触发自动链路。 "auto_triview": False, @@ -1440,6 +2029,9 @@ def generate_base_asset(*, project, user, kind: str, prompt: str, label: str = " model_config=model_config, request_payload=payload, ) + # 真实平台成本由每条 AIModelAttempt 按实际调用模型累加,避免 Fallback 后仍记默认模型旧成本。 + task.base_cost = Decimal("0") + task.save(update_fields=["base_cost", "updated_at"]) generate_base_asset_task.delay(str(task.id)) return task @@ -1460,28 +2052,47 @@ def run_base_asset_task(*, task_id: str) -> None: use_edit = bool(payload.get("use_edit")) ref_url = str(payload.get("reference_image") or "") model_config = task.model_config - provider = get_image_provider(model_config) + use_model_routing = bool(payload.get("model_routing_v1")) + provider = None if use_model_routing else get_image_provider(model_config) reservation = task.credit_reservation try: + if use_edit and ref_url: + # 商品三视图默认横;角色立绘重跑默认竖,与原调用尺寸保持一致。 + if kind == BaseAssetGroup.Kind.PERSON: + edit_size = prompt_ratio_size("person_portrait", "1024x1536") + else: + edit_size = prompt_ratio_size("product_triview", "1536x1024") + generate_size = None + else: + edit_size = None + # 场景默认横、人物立绘默认竖;比例仍由现有提示词配置决定。 + if kind == BaseAssetGroup.Kind.SCENE: + generate_size = prompt_ratio_size("scene", "1536x1024") + else: + generate_size = prompt_ratio_size("person_portrait", "1024x1536") + def _make(): if use_edit and ref_url: - # 商品三视图默认横;角色立绘重跑(image_edit 参考当前立绘)默认竖,与文生图立绘一致。 - if kind == BaseAssetGroup.Kind.PERSON: - size = prompt_ratio_size("person_portrait", "1024x1536") - else: - size = prompt_ratio_size("product_triview", "1536x1024") - resp = provider.image_edit(model=model_config.name, prompt=prompt, images=[ref_url], size=size) + resp = provider.image_edit(model=model_config.name, prompt=prompt, images=[ref_url], size=edit_size) else: - # 场景默认横、人物立绘默认竖;两者比例都可在 admin 改 - if kind == BaseAssetGroup.Kind.SCENE: - gen_size = prompt_ratio_size("scene", "1536x1024") - else: - gen_size = prompt_ratio_size("person_portrait", "1024x1536") - resp = provider.image_generation(model=model_config.name, endpoint=model_config.endpoint, prompt=prompt, size=gen_size) + resp = provider.image_generation(model=model_config.name, endpoint=model_config.endpoint, prompt=prompt, size=generate_size) return resp, provider.extract_first_media_url(resp) - # 中转站偶发抖动(网络/5xx/限流/空返回)按退避重试,避免「一次失败就生成不出来」;永久错误仍立即退费 - response, media = _run_image_with_retry(_make) + if use_model_routing: + routed = execute_routed_image_request( + task=task, + primary_model=model_config, + prompt=prompt, + reference_images=[ref_url] if use_edit and ref_url else [], + edit_size=edit_size, + direct_size=edit_size, + generate_size=generate_size, + request_summary={"base_asset_kind": kind}, + ) + response, media = routed.value + else: + # 存量未迁移任务继续沿用旧局部重试,避免部署切换期间改变已排队任务行为。 + response, media = _run_image_with_retry(_make) with transaction.atomic(): task.status = AITask.Status.SUCCEEDED task.response_payload = response @@ -1557,14 +2168,21 @@ def generate_person_triview(*, project, user, portrait_asset) -> AITask: model_config = get_default_model(ModelConfig.Capability.IMAGE) if model_config is None: raise ValueError("no active image model configured") - provider = get_image_provider(model_config) - if not hasattr(provider, "image_edit"): - raise ValueError(f"当前图像模型 {model_config.provider.name}:{model_config.name} 不支持参考图三视图(image_edit)") ref_url = _asset_preview_url(portrait_asset) # 人物三视图提示词:正文可在 admin「提示词」页改(无占位符) tri_prompt = render_prompt("person_triview", THREE_VIEW_PROMPT) - payload = {"model": model_config.name, "prompt": tri_prompt, "kind": "person", "triview_of": asset_key, "reference_image": ref_url} + payload = { + "model": model_config.name, + "prompt": tri_prompt, + "kind": "person", + "triview_of": asset_key, + "reference_image": ref_url, + "model_routing_v1": True, + } task = create_ai_task(project=project, user=user, task_type=AITask.Type.PERSON_IMAGE, model_config=model_config, request_payload=payload) + # 真实平台成本由每条 AIModelAttempt 按实际模型累加,避免 Fallback 后仍记默认模型旧成本。 + task.base_cost = Decimal("0") + task.save(update_fields=["base_cost", "updated_at"]) generate_triview_task.delay(str(task.id)) return task @@ -1599,9 +2217,6 @@ def generate_model_triview(*, model, user) -> tuple[AITask, bool]: if model.portrait_asset_id is None: raise ValueError("请先设置模特形象图") model_config, quote = quote_model_triview(model=model) - provider = get_image_provider(model_config) - if not hasattr(provider, "image_edit"): - raise ValueError(f"当前图像模型 {model_config.provider.name}:{model_config.name} 不支持参考图三视图(image_edit)") ref_url = _asset_preview_url(model.portrait_asset) if not ref_url: raise ValueError("模特形象图不可用,请先重新上传") @@ -1629,6 +2244,7 @@ def generate_model_triview(*, model, user) -> tuple[AITask, bool]: "portrait_asset_id": str(locked.portrait_asset_id), "reference_image": ref_url, "price_points": str(quote.points), + "model_routing_v1": True, } task = AITask.objects.create( team=locked.team, @@ -1640,7 +2256,8 @@ def generate_model_triview(*, model, user) -> tuple[AITask, bool]: idempotency_key=f"model_triview:{locked.id}:{uuid.uuid4()}", request_payload=payload, estimated_cost=quote.points, - base_cost=quote.base_cost_yuan, + # 路由任务的平台成本由每条 AIModelAttempt 按实际模型累加。 + base_cost=Decimal("0"), ) reserve_credit(team=locked.team, user=user, task=task, amount=quote.points) task.status = AITask.Status.RESERVED @@ -1672,14 +2289,16 @@ def run_model_triview_task(*, task_id: str) -> None: portrait_asset_id = str(payload.get("portrait_asset_id") or "") ref_url = str(payload.get("reference_image") or "") prompt = str(payload.get("prompt") or "") - provider = get_image_provider(task.model_config) + use_model_routing = bool(payload.get("model_routing_v1")) + provider = None if use_model_routing else get_image_provider(task.model_config) reservation = task.credit_reservation try: if not model_id or not portrait_asset_id or not ref_url: raise ValueError("模特三视图任务参数不完整") + size = prompt_ratio_size("person_triview", "1536x864") + def _make(): - size = prompt_ratio_size("person_triview", "1536x1024") response = provider.image_edit( model=task.model_config.name, prompt=prompt, @@ -1688,7 +2307,20 @@ def run_model_triview_task(*, task_id: str) -> None: ) return response, provider.extract_first_media_url(response) - response, media = _run_image_with_retry(_make) + if use_model_routing: + routed = execute_routed_image_request( + task=task, + primary_model=task.model_config, + prompt=prompt, + reference_images=[ref_url], + aspect_ratio="16:9", + edit_size=size, + direct_size=size, + request_summary={"model_triview": True}, + ) + response, media = routed.value + else: + response, media = _run_image_with_retry(_make) with transaction.atomic(): locked_model = ( Model.objects.select_for_update() @@ -1780,16 +2412,31 @@ def run_triview_task(*, task_id: str) -> None: ref_url = str(payload.get("reference_image") or "") prompt = str(payload.get("prompt") or THREE_VIEW_PROMPT) model_config = task.model_config - provider = get_image_provider(model_config) + use_model_routing = bool(payload.get("model_routing_v1")) + provider = None if use_model_routing else get_image_provider(model_config) reservation = task.credit_reservation try: + tri_size = prompt_ratio_size("person_triview", "1536x864") + def _make(): - tri_size = prompt_ratio_size("person_triview", "1536x1024") resp = provider.image_edit(model=model_config.name, prompt=prompt, images=[ref_url], size=tri_size) return resp, provider.extract_first_media_url(resp) - # 同基础资产:瞬时抖动重试,永久错误立即退费报错 - response, media = _run_image_with_retry(_make) + if use_model_routing: + routed = execute_routed_image_request( + task=task, + primary_model=model_config, + prompt=prompt, + reference_images=[ref_url], + aspect_ratio="16:9", + edit_size=tri_size, + direct_size=tri_size, + request_summary={"project_person_triview": True}, + ) + response, media = routed.value + else: + # 存量未迁移任务继续沿用旧局部重试,避免部署切换期间改变已排队任务行为。 + response, media = _run_image_with_retry(_make) with transaction.atomic(): task.status = AITask.Status.SUCCEEDED task.response_payload = response @@ -2259,24 +2906,53 @@ def _storyboard_shot_worker(task_id, shot_id, user_id) -> None: task.status = AITask.Status.SUBMITTED task.save(update_fields=["status", "updated_at"]) try: - provider = get_image_provider(model_config) + use_model_routing = bool((task.request_payload or {}).get("model_routing_v1")) + provider = None if use_model_routing else get_image_provider(model_config) refs = _storyboard_reference_images(project, segment) if segment is not None else [] ref_urls = [r["url"] for r in refs] - if ref_urls and hasattr(provider, "image_edit"): + if ref_urls and (use_model_routing or hasattr(provider, "image_edit")): # gpt-image-2 多图参考:必须用 refs 版提示词(点名「参考图N=角色/场景/商品」+锁脸锁商品) frame_prompt = build_storyboard_frame_prompt_refs(project, segment, refs, extra_prompt) - response = _call_image_with_retry( - lambda: provider.image_edit(model=model_config.name, prompt=frame_prompt, images=ref_urls, size="1024x1536") - ) else: frame_prompt = ( build_storyboard_frame_prompt(project, segment, extra_prompt) if segment is not None else (task.request_payload.get("prompt") or "") ) - response = _call_image_with_retry( - lambda: provider.image_generation(model=model_config.name, endpoint=model_config.endpoint, prompt=frame_prompt) + + if use_model_routing: + routed = execute_routed_image_request( + task=task, + primary_model=model_config, + prompt=frame_prompt, + reference_images=ref_urls, + aspect_ratio="9:16", + edit_size="1024x1536", + direct_size="1024x1536", + request_summary={ + "storyboard_shot": str(shot.id), + "storyboard_sort_order": shot.sort_order, + }, ) - media = provider.extract_first_media_url(response) + response, media = routed.value + else: + if ref_urls and hasattr(provider, "image_edit"): + response = _call_image_with_retry( + lambda: provider.image_edit( + model=model_config.name, + prompt=frame_prompt, + images=ref_urls, + size="1024x1536", + ) + ) + else: + response = _call_image_with_retry( + lambda: provider.image_generation( + model=model_config.name, + endpoint=model_config.endpoint, + prompt=frame_prompt, + ) + ) + media = provider.extract_first_media_url(response) asset = _store_generated_media( team=project.team, user=user, project=project, task=task, media=media, name=f"{project.name}-storyboard-{shot.sort_order + 1}", @@ -2367,8 +3043,12 @@ def poll_storyboard(*, project, user) -> dict: "model": model_config.name, "endpoint": model_config.endpoint, "prompt": build_storyboard_frame_prompt(project, segment, extra_prompt) if segment is not None else "", "storyboard_shot": str(shot.id), + "model_routing_v1": True, }, ) + # 真实平台成本由每条 AIModelAttempt 按实际模型累加,避免 Fallback 后仍记默认模型旧成本。 + task.base_cost = Decimal("0") + task.save(update_fields=["base_cost", "updated_at"]) StoryboardShot.objects.filter(id=shot.id).update(status=StoryboardShot.Status.RUNNING, updated_at=timezone.now()) threading.Thread( target=_storyboard_shot_worker, args=(str(task.id), str(shot.id), str(user.id)), daemon=True @@ -2550,26 +3230,41 @@ def submit_video_segment(*, video_segment: VideoSegment, user, prompt: str) -> V "price_multiplier": quote.meta.get("price_multiplier", "1"), "video_segment_id": str(video_segment.id), "reference_images": reference_images, + "model_routing_v1": True, }, ) + # 提交尝试的实际平台成本由 AIModelAttempt 累加;成片后再用实际模型的 usage true-up 覆盖。 + task.base_cost = Decimal("0") + task.save(update_fields=["base_cost", "updated_at"]) try: - provider = build_provider(model_config) - # 不再静默退文生兜底:火山报错(如人脸需走素材库的 InputImageSensitiveContentDetected)直接抛, - # 由下方 except 保存原始错误供排障,并将安全提示写给普通用户。 - 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=final_prompt, duration=video_segment.target_duration_seconds, ratio="9:16", resolution="720p", - reference_images=reference_images or None, + reference_images=reference_images, + request_summary={"video_segment_id": str(video_segment.id)}, ) - task.provider_task_id = str(response.get("id") or response.get("task_id") or "") + response, provider_task_id = routed.value + 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", + ] + ) video_segment.status = VideoSegment.Status.RUNNING video_segment.save(update_fields=["status", "updated_at"]) return None @@ -2620,8 +3315,43 @@ def poll_video_segment(*, video_segment: VideoSegment, user) -> VideoSegmentVers if ai_task.status in (AITask.Status.FAILED, AITask.Status.CANCELLED): return None - provider = build_provider(ai_task.model_config) - response = provider.poll_video_task(endpoint=ai_task.model_config.endpoint, provider_task_id=ai_task.provider_task_id) + # Fallback 只发生在提交阶段;拿到远端任务 ID 后必须固定到实际提交成功的模型和 Provider 轮询。 + submit_attempt = ( + ai_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 ai_task.model_config + from apps.ai.routing_policy import load_model_routing_policy + + video_policy = load_model_routing_policy().video + if ai_task.submitted_at and (timezone.now() - ai_task.submitted_at).total_seconds() >= video_policy.generation_timeout: + timeout_message = "视频生成超过配置的等待总时限" + with transaction.atomic(): + locked_task = AITask.objects.select_for_update().get(id=ai_task.id) + if locked_task.status in (AITask.Status.SUCCEEDED, AITask.Status.FAILED, AITask.Status.CANCELLED): + return video_segment.versions.filter(task=locked_task).order_by("-created_at").first() + locked_task.status = AITask.Status.FAILED + locked_task.error_message = timeout_message + locked_task.completed_at = timezone.now() + locked_task.save(update_fields=["status", "error_message", "completed_at", "updated_at"]) + release_credit(reservation=locked_task.credit_reservation, reason=timeout_message) + video_segment.status = VideoSegment.Status.FAILED + video_segment.error_message = classify_generation_error( + TimeoutError(timeout_message), + operation="video_generate", + reference_id=str(locked_task.id), + ).fallback_message + video_segment.save(update_fields=["status", "error_message", "updated_at"]) + return None + + provider = get_video_provider(actual_model) + response = provider.poll_video_task( + endpoint=actual_model.endpoint, + provider_task_id=ai_task.provider_task_id, + timeout=video_policy.poll_request_timeout, + ) remote_status = response.get("status") if remote_status in {"queued", "running", "processing"}: # 仍在生成:只在状态首次进入 POLLING 时落一次库。旧实现每次 poll(5s 一次)都把完整 @@ -2687,7 +3417,7 @@ def poll_video_segment(*, video_segment: VideoSegment, user) -> VideoSegmentVers usage_tokens = 0 if usage_tokens > 0: settle = quote_video_actual( - locked_task.model_config, + actual_model, tokens=usage_tokens, with_video_ref=False, resolution=str(payload.get("resolution") or "720p"), @@ -2792,10 +3522,19 @@ def enqueue_standalone_images(*, team, user, prompt: str, mode: str = "image", c task_type = _STANDALONE_TASK_TYPE.get(mode, AITask.Type.PRODUCT_IMAGE) count = max(1, min(int(count or 1), 12)) ref_ids = [str(r) for r in (reference_image_ids or []) if r] - # 普通图片创作带用户上传参考图时,当前模型不支持 image_edit 就沿用既有保护切到 gpt-image。 - # 模特上身图例外:商品图和模特图由专用 Worker 作为多参考图传入 Seedream/GPT,必须尊重用户自选模型; - # 即使旧客户端额外带了 reference_image_ids,也不得在任务创建阶段静默覆盖 image_model。 - if ref_ids and mode != "model" and not hasattr(build_provider(model_config), "image_edit"): + # 已迁移范围: + # 1) 普通图片创作 + 单张,覆盖无参考图、单参考图和多参考图; + # 2) 绑定商品的模特上身图,每张仍是独立 AITask、独立尝试链; + # 3) 绑定商品的平台套图,每个平台、每张图仍沿用原有独立任务与批次关系。 + # 新建模特候选(mode=model 但无 product_id)等后续入口继续保持原流程。 + use_model_routing = ( + (mode == "image" and count == 1) + or (mode == "model" and bool(product_id)) + or (mode == "cover" and bool(product_id)) + ) and not reference_product + # 尚未迁移的非图片创作模式保留既有兼容保护;图片创作必须尊重用户选择的主模型, + # Seedream 可通过 image_generation(image=...) 处理参考图,失败后再由统一路由动态切换。 + if ref_ids and mode not in {"model", "image"} and not hasattr(build_provider(model_config), "image_edit"): alt = resolve_image_model("gpt-image") if alt is not None: model_config = alt @@ -2880,6 +3619,8 @@ def enqueue_standalone_images(*, team, user, prompt: str, mode: str = "image", c for index in range(count): quote = quote_flat(model_config, team=team) request_payload = {"model": model_config.name, "endpoint": model_config.endpoint, "prompt": prompt, "mode": mode, "index": index, "product_id": str(product_id) if product_id else None, "reference_product": bool(reference_product), "model_id": str(model_id) if model_id else None, "model_entity_id": str(model_entity_id) if model_entity_id else None, "batch_id": batch_id, "ratio": str(ratio) if ratio else None, "reference_image_ids": ref_ids, "platform_id": platform_key or None, "platform_name": platform_name or None} + if use_model_routing: + request_payload["model_routing_v1"] = True if tryon_classification is not None: request_payload["tryon_classification"] = dict(tryon_classification) if tryon_trouser_facts is not None: @@ -2911,7 +3652,8 @@ def enqueue_standalone_images(*, team, user, prompt: str, mode: str = "image", c idempotency_key=f"standalone-image:{team.id}:{uuid.uuid4()}", request_payload=request_payload, estimated_cost=quote.points, - base_cost=quote.base_cost_yuan, + # 路由入口的平台成本改由每条 AIModelAttempt 累加,避免候选切换后仍记主模型旧成本。 + base_cost=Decimal("0") if use_model_routing else quote.base_cost_yuan, ) # 预留额度若余额不足会抛 ValueError,在同步的 Web 请求里立刻反馈给前端(不会先建半套任务) reserve_credit(team=team, user=user, task=task, amount=quote.points) @@ -2945,7 +3687,13 @@ def run_standalone_image_task(*, task_id: str) -> None: else: category = _STANDALONE_CATEGORY.get(mode, Asset.Category.UNCATEGORIZED) model_config = task.model_config - provider = get_image_provider(model_config) + use_model_routing = ( + bool(payload.get("model_routing_v1")) + and mode in {"image", "model", "cover"} + and not payload.get("reference_product") + ) + # 路由入口必须在每次尝试时按真实候选重新构造 Provider;提前构造会让主模型配置错误绕过尝试日志。 + provider = None if use_model_routing else get_image_provider(model_config) reservation = task.credit_reservation # 出图策略(都优先 image_edit 锁真实素材,模型不支持/无素材才回落纯文生图): # · 模特上身图(mode=model):参考图1=商品真实主图 + 参考图2=选中模特 → 生成「该模特用该商品」效果图; @@ -2957,7 +3705,7 @@ def run_standalone_image_task(*, task_id: str) -> None: from apps.products.models import Product product = Product.objects.filter(id=product_id, team=team).first() - can_edit = hasattr(provider, "image_edit") + can_edit = hasattr(provider, "image_edit") if provider is not None else False model_url = "" if payload.get("model_id"): model_asset = Asset.objects.filter(id=payload.get("model_id")).first() @@ -3066,7 +3814,26 @@ def run_standalone_image_task(*, task_id: str) -> None: payload["tryon_prompt"] = tryon_prompt_trace task.request_payload = payload task.save(update_fields=["request_payload", "updated_at"]) - if use_edit and can_edit: + # 纯文生图同批多张要不同构图,否则雷同(PMC#24);index>0 追加换版式指令。 + # 当前统一路由只接入单张图片创作,但这里仍保留原提示词构造供未迁移批量旧路径复用。 + gen_prompt = prompt + if index > 0: + variation = _FREE_VARIATIONS[index % len(_FREE_VARIATIONS)] + gen_prompt = f"{prompt}。{variation},不要在画面中添加任何文字、卖点、标题或价签。" + + if use_model_routing: + call_prompt = edit_prompt if edit_images else gen_prompt + routed = execute_routed_image_request( + task=task, + primary_model=model_config, + prompt=call_prompt, + reference_images=edit_images, + aspect_ratio=output_ratio or None, + edit_size=_ratio_to_image_size(output_ratio), + direct_size=_ratio_to_volcano_size(output_ratio), + ) + response, media = routed.value + elif use_edit and can_edit: # gpt-image 等支持 image_edit(多图参考编辑接口) if payload.get("reference_product"): size = "1536x1024" # 三视图固定横向 @@ -3079,15 +3846,9 @@ def run_standalone_image_task(*, task_id: str) -> None: vsize = "2304x1728" if payload.get("reference_product") else _ratio_to_volcano_size(output_ratio) response = provider.image_generation(model=model_config.name, endpoint=model_config.endpoint, prompt=edit_prompt, image=edit_images, size=vsize) else: - # 纯文生图同批多张要不同构图,否则雷同(PMC#24);index>0 追加换版式指令。 - # 自由模式只换构图/视角/光线、绝不加文字卖点(_FREE_VARIATIONS); - # 套图卖点版式归 mode=cover,不在这条裸文生图分支里。 - gen_prompt = prompt - if index > 0: - variation = _FREE_VARIATIONS[index % len(_FREE_VARIATIONS)] - gen_prompt = f"{prompt}。{variation},不要在画面中添加任何文字、卖点、标题或价签。" response = provider.image_generation(model=model_config.name, endpoint=model_config.endpoint, prompt=gen_prompt) - media = provider.extract_first_media_url(response) + if not use_model_routing: + media = provider.extract_first_media_url(response) with transaction.atomic(): task.status = AITask.Status.SUCCEEDED task.response_payload = response @@ -3168,16 +3929,22 @@ def synthesize_project_voiceover(*, project, user, items: list[dict], voice_type texts.append((idx, j, piece)) if not texts: raise ValueError("没有可配音的旁白文本") - provider = VolcanoTtsProvider() - if not provider.configured: - raise TtsNotConfigured( - "语音合成未配置:请在后端环境变量设置 VOLC_TTS_APPID 和 VOLC_TTS_ACCESS_TOKEN" - "(火山引擎控制台 → 语音技术 → 语音合成大模型 → 创建应用)" - ) voice_type = voice_type or DEFAULT_VOICEOVER_VOICE model_config = get_default_model(ModelConfig.Capability.AUDIO) if model_config is None: raise ValueError("no active audio model configured") + primary_provider = get_audio_provider(model_config) + # 保持当前直连未配置时“不建任务、不预扣”的旧体验;管理员将该模型的向外 Fallback 开关 + # 打开后,则允许统一执行器把配置故障路由到已启用的 OpenAI 兼容音频候选。 + if ( + hasattr(primary_provider, "configured") + and not primary_provider.configured + and not model_allows_fallback(model_config) + ): + raise TtsNotConfigured( + "语音合成未配置:请在后端环境变量设置 VOLC_TTS_APPID 和 VOLC_TTS_ACCESS_TOKEN" + "(火山引擎控制台 → 语音技术 → 语音合成大模型 → 创建应用)" + ) # 配音按字符数阶梯计价(默认每 500 字 10 积分,不足按 500):长短脚本不再同价 from apps.billing.pricing import quote_voiceover @@ -3194,13 +3961,28 @@ def synthesize_project_voiceover(*, project, user, items: list[dict], voice_type "speed_ratio": float(speed_ratio or 1.0), "char_count": char_count, "items": [{"index": idx, "cue": j, "text": text} for idx, j, text in texts], + "model_routing_v1": True, }, ) reservation = task.credit_reservation + 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"]) synthesized = [] for idx, j, text in texts: - audio, duration_ms = provider.synthesize(text=text, voice_type=voice_type, speed_ratio=speed_ratio, uid=str(user.id)) + routed = execute_routed_audio_request( + task=task, + primary_model=model_config, + text=text, + public_voice=voice_type, + speed_ratio=speed_ratio, + user_id=str(user.id), + request_summary={"segment_index": idx, "cue_index": j}, + ) + audio, duration_ms = routed.value synthesized.append((idx, j, text, audio, duration_ms)) with transaction.atomic(): task.status = AITask.Status.SUCCEEDED diff --git a/core/backend/apps/ai/tasks.py b/core/backend/apps/ai/tasks.py index f513d62..5490f1b 100644 --- a/core/backend/apps/ai/tasks.py +++ b/core/backend/apps/ai/tasks.py @@ -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 diff --git a/core/backend/apps/ai/test_base_asset_routing.py b/core/backend/apps/ai/test_base_asset_routing.py new file mode 100644 index 0000000..bdc543c --- /dev/null +++ b/core/backend/apps/ai/test_base_asset_routing.py @@ -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) diff --git a/core/backend/apps/ai/test_entity_extraction_routing.py b/core/backend/apps/ai/test_entity_extraction_routing.py new file mode 100644 index 0000000..58e2a73 --- /dev/null +++ b/core/backend/apps/ai/test_entity_extraction_routing.py @@ -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) diff --git a/core/backend/apps/ai/test_free_video_routing.py b/core/backend/apps/ai/test_free_video_routing.py new file mode 100644 index 0000000..d720d01 --- /dev/null +++ b/core/backend/apps/ai/test_free_video_routing.py @@ -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")) diff --git a/core/backend/apps/ai/test_model_routing.py b/core/backend/apps/ai/test_model_routing.py new file mode 100644 index 0000000..36191d2 --- /dev/null +++ b/core/backend/apps/ai/test_model_routing.py @@ -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) diff --git a/core/backend/apps/ai/test_model_triview_routing.py b/core/backend/apps/ai/test_model_triview_routing.py new file mode 100644 index 0000000..0998c86 --- /dev/null +++ b/core/backend/apps/ai/test_model_triview_routing.py @@ -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) diff --git a/core/backend/apps/ai/test_project_triview_routing.py b/core/backend/apps/ai/test_project_triview_routing.py new file mode 100644 index 0000000..dc0d6f0 --- /dev/null +++ b/core/backend/apps/ai/test_project_triview_routing.py @@ -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) diff --git a/core/backend/apps/ai/test_routing_executor.py b/core/backend/apps/ai/test_routing_executor.py new file mode 100644 index 0000000..5165e5d --- /dev/null +++ b/core/backend/apps/ai/test_routing_executor.py @@ -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")) diff --git a/core/backend/apps/ai/test_routing_policy.py b/core/backend/apps/ai/test_routing_policy.py new file mode 100644 index 0000000..73cfb26 --- /dev/null +++ b/core/backend/apps/ai/test_routing_policy.py @@ -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) diff --git a/core/backend/apps/ai/test_script_entity_sync.py b/core/backend/apps/ai/test_script_entity_sync.py new file mode 100644 index 0000000..218d4f4 --- /dev/null +++ b/core/backend/apps/ai/test_script_entity_sync.py @@ -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, + ) diff --git a/core/backend/apps/ai/test_script_optimization_routing.py b/core/backend/apps/ai/test_script_optimization_routing.py new file mode 100644 index 0000000..7aa6d8c --- /dev/null +++ b/core/backend/apps/ai/test_script_optimization_routing.py @@ -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")) diff --git a/core/backend/apps/ai/test_script_stream_routing.py b/core/backend/apps/ai/test_script_stream_routing.py new file mode 100644 index 0000000..e4d767a --- /dev/null +++ b/core/backend/apps/ai/test_script_stream_routing.py @@ -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) diff --git a/core/backend/apps/ai/test_standalone_image_routing.py b/core/backend/apps/ai/test_standalone_image_routing.py new file mode 100644 index 0000000..ad5de39 --- /dev/null +++ b/core/backend/apps/ai/test_standalone_image_routing.py @@ -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) diff --git a/core/backend/apps/ai/test_storyboard_routing.py b/core/backend/apps/ai/test_storyboard_routing.py new file mode 100644 index 0000000..10d0782 --- /dev/null +++ b/core/backend/apps/ai/test_storyboard_routing.py @@ -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)) diff --git a/core/backend/apps/ai/test_video_segment_routing.py b/core/backend/apps/ai/test_video_segment_routing.py new file mode 100644 index 0000000..cce82ca --- /dev/null +++ b/core/backend/apps/ai/test_video_segment_routing.py @@ -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")) diff --git a/core/backend/apps/ai/test_voiceover_routing.py b/core/backend/apps/ai/test_voiceover_routing.py new file mode 100644 index 0000000..1d1123f --- /dev/null +++ b/core/backend/apps/ai/test_voiceover_routing.py @@ -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) diff --git a/core/backend/apps/ai/tests.py b/core/backend/apps/ai/tests.py index d938a96..3cab6cd 100644 --- a/core/backend/apps/ai/tests.py +++ b/core/backend/apps/ai/tests.py @@ -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 向两类模型传商品/人物参考图,不能沿用自由创作的自动换模逻辑。""" diff --git a/core/backend/apps/common/management/commands/bootstrap_volcano_models.py b/core/backend/apps/common/management/commands/bootstrap_volcano_models.py index 4baf569..f4f8df5 100644 --- a/core/backend/apps/common/management/commands/bootstrap_volcano_models.py +++ b/core/backend/apps/common/management/commands/bootstrap_volcano_models.py @@ -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: diff --git a/core/backend/apps/common/management/commands/bootstrap_yunqi_models.py b/core/backend/apps/common/management/commands/bootstrap_yunqi_models.py index e211066..053696b 100644 --- a/core/backend/apps/common/management/commands/bootstrap_yunqi_models.py +++ b/core/backend/apps/common/management/commands/bootstrap_yunqi_models.py @@ -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: diff --git a/core/frontend/src/admin-page.css b/core/frontend/src/admin-page.css index adf119d..ea242ad 100644 --- a/core/frontend/src/admin-page.css +++ b/core/frontend/src/admin-page.css @@ -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 { diff --git a/core/frontend/src/routes/admin/admin-tasks.tsx b/core/frontend/src/routes/admin/admin-tasks.tsx index cd321fc..e9e6c76 100644 --- a/core/frontend/src/routes/admin/admin-tasks.tsx +++ b/core/frontend/src/routes/admin/admin-tasks.tsx @@ -30,6 +30,12 @@ function statusPill(status: string) { return 进行中; } +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([]); const [count, setCount] = useState(0); @@ -171,13 +177,55 @@ export function AdminTasksPage({ notify }: { notify: Notify }) {

加载中…

) : ( <> + {(() => { + const attempts = detail.attempts || []; + const finalAttempt = [...attempts].reverse().find((attempt) => attempt.status === "succeeded"); + const fallbackUsed = attempts.some((attempt) => attempt.is_fallback); + return (
状态{statusPill(detail.status)}
团队{detail.team_name || "—"}
模型{detail.model_name || "—"}
计价{pts(detail.estimated_cost)} → {pts(detail.actual_cost)} 积分{detail.cost_anomaly ? " ⚠" : ""}
平台成本{Number(detail.base_cost || 0) > 0 ? `¥${detail.base_cost} · 毛利 ${detail.margin_yuan != null ? `¥${detail.margin_yuan}` : "—"}` : "未知"}
+
模型调用{attempts.length ? `${attempts.length} 次 · ${fallbackUsed ? "发生 Fallback" : "未切换"}` : "旧任务 · 无尝试链"}
+ {finalAttempt &&
最终成功模型{finalAttempt.provider_display_name || finalAttempt.provider_name} / {finalAttempt.model_display_name || finalAttempt.model_name}
}
+ ); + })()} + {detail.attempts?.length > 0 && ( + <> +
模型调用链 · {detail.attempts.length} 次
+
+ {detail.attempts.map((attempt) => ( +
+
+ // {String(attempt.sequence).padStart(2, "0")} + {statusPill(attempt.status)} + {attemptKind(attempt)} +
+
+ {attempt.provider_display_name || attempt.provider_name} + / + {attempt.model_display_name || attempt.model_name} +
+
+ {attempt.operation} + {attempt.duration_ms == null ? "耗时 —" : `耗时 ${attempt.duration_ms} ms`} + {Number(attempt.platform_cost || 0) > 0 ? `平台成本 ¥${attempt.platform_cost}` : "平台成本未知"} + {attempt.provider_task_id && Provider ID {attempt.provider_task_id}} +
+ {(attempt.error_type || attempt.raw_error) && ( +
+ [{attempt.error_type || "unknown"}{attempt.provider_error_code ? ` · ${attempt.provider_error_code}` : ""}] + {attempt.raw_error || attempt.safe_error_summary} +
+ )} +
+ ))} +
+ + )} {detail.error_message && ( <>
错误
diff --git a/core/frontend/src/types.ts b/core/frontend/src/types.ts index 89ce275..f52da1f 100644 --- a/core/frontend/src/types.ts +++ b/core/frontend/src/types.ts @@ -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; + platform_cost: string; + request_summary: Record; + response_summary: Record; +}; 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; diff --git a/core/qa/model-routing-browser-qa.py b/core/qa/model-routing-browser-qa.py new file mode 100644 index 0000000..d234226 --- /dev/null +++ b/core/qa/model-routing-browser-qa.py @@ -0,0 +1,1651 @@ +"""为模型路由浏览器验收创建可删除的 QA 任务。 + +只允许在 DEBUG 环境且显式设置 MODEL_ROUTING_BROWSER_QA=1 时运行。脚本不会调用真实 +Provider 或对象存储,也不会使用现有团队积分;它创建不可登录的临时用户和团队,真实走 +AITask、预留积分、统一路由执行器、尝试日志及最终结算/释放流程。 + +用法(从 core/backend 目录运行): + MODEL_ROUTING_BROWSER_QA=1 python ../../core/qa/model-routing-browser-qa.py setup + MODEL_ROUTING_BROWSER_QA=1 python ../../core/qa/model-routing-browser-qa.py setup-reference + MODEL_ROUTING_BROWSER_QA=1 python ../../core/qa/model-routing-browser-qa.py setup-tryon + MODEL_ROUTING_BROWSER_QA=1 python ../../core/qa/model-routing-browser-qa.py setup-platform + MODEL_ROUTING_BROWSER_QA=1 python ../../core/qa/model-routing-browser-qa.py setup-base + MODEL_ROUTING_BROWSER_QA=1 python ../../core/qa/model-routing-browser-qa.py setup-model-triview + MODEL_ROUTING_BROWSER_QA=1 python ../../core/qa/model-routing-browser-qa.py setup-project-triview + MODEL_ROUTING_BROWSER_QA=1 python ../../core/qa/model-routing-browser-qa.py setup-storyboard + MODEL_ROUTING_BROWSER_QA=1 python ../../core/qa/model-routing-browser-qa.py setup-entity + MODEL_ROUTING_BROWSER_QA=1 python ../../core/qa/model-routing-browser-qa.py setup-script + MODEL_ROUTING_BROWSER_QA=1 python ../../core/qa/model-routing-browser-qa.py setup-voice + MODEL_ROUTING_BROWSER_QA=1 python ../../core/qa/model-routing-browser-qa.py setup-video + MODEL_ROUTING_BROWSER_QA=1 python ../../core/qa/model-routing-browser-qa.py setup-free-video + MODEL_ROUTING_BROWSER_QA=1 python ../../core/qa/model-routing-browser-qa.py cleanup +""" + +from __future__ import annotations + +import argparse +from copy import deepcopy +from decimal import Decimal +from io import BytesIO +import json +import os +from pathlib import Path +import sys +from unittest.mock import Mock, patch + + +BACKEND_DIR = Path(__file__).resolve().parents[1] / "backend" +sys.path.insert(0, str(BACKEND_DIR)) +os.environ.setdefault("DJANGO_SETTINGS_MODULE", "airshelf.settings.development") + +import django # noqa: E402 + +django.setup() + +import requests # noqa: E402 +from django.conf import settings # noqa: E402 +from django.test import override_settings # noqa: E402 + +from apps.accounts.models import Team, TeamMember, User # noqa: E402 +from apps.ai.models import AITask, ModelConfig, ModelProvider # noqa: E402 +from apps.ai.script_agent import regenerate_segment_via_agent, stream_script_agent # noqa: E402 +from apps.ai.services import ( # noqa: E402 + _storyboard_shot_worker, + enqueue_standalone_images, + generate_base_asset, + generate_model_triview, + generate_person_triview, + run_extract_entities_task, + run_base_asset_task, + run_model_triview_task, + run_standalone_image_task, + run_triview_task, + poll_storyboard, + poll_video_segment, + submit_video_segment, + submit_extract_entities, + synthesize_project_voiceover, +) +from apps.assets.models import Asset, AssetFile, Model # noqa: E402 +from apps.assets.storage import StoredObject # noqa: E402 +from apps.billing.models import CreditAccount, CreditLedger # noqa: E402 +from apps.products.models import Product # noqa: E402 +from apps.projects.models import ( # noqa: E402 + BaseAssetGroup, + Project, + ScriptSegment, + ScriptVersion, + StoryboardShot, + VideoSegment, +) + + +QA_USERNAME = "__routing_browser_qa__" +QA_TEAM_NAME = "__Routing Browser QA__" +QA_PROVIDER_NAMES = ( + "qa-routing-primary", + "qa-routing-fallback-1", + "qa-routing-fallback-2", +) + + +def _assert_safe() -> None: + if not settings.DEBUG or os.environ.get("MODEL_ROUTING_BROWSER_QA") != "1": + raise SystemExit( + "拒绝运行:仅允许在 DEBUG 环境,并需显式设置 MODEL_ROUTING_BROWSER_QA=1。" + ) + + +def cleanup() -> dict[str, int]: + tasks = AITask.objects.filter(team__name=QA_TEAM_NAME).count() + teams = Team.objects.filter(name=QA_TEAM_NAME).count() + users = User.objects.filter(username=QA_USERNAME).count() + providers = ModelProvider.objects.filter(name__in=QA_PROVIDER_NAMES).count() + + # 项目通过 PROTECT 引用商品;先删精确 QA 团队下的项目,再让团队级联清理商品与其余数据。 + Project.objects.filter(team__name=QA_TEAM_NAME).delete() + Team.objects.filter(name=QA_TEAM_NAME).delete() + ModelProvider.objects.filter(name__in=QA_PROVIDER_NAMES).delete() + User.objects.filter(username=QA_USERNAME).delete() + return {"tasks": tasks, "teams": teams, "users": users, "providers": providers} + + +def _image_metadata(*, priority: int) -> dict: + return { + "routing": { + "fallback_on_failure": True, + "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": "0.10"}, + "qa": {"browser_routing_priority": priority}, + } + + +def _create_model(provider_name: str, priority: int, model_name: str) -> ModelConfig: + provider = ModelProvider.objects.create( + name=provider_name, + display_name=f"QA Provider {priority}", + status=ModelProvider.Status.ACTIVE, + base_url="http://127.0.0.1.invalid/v1", + api_key="", + metadata={"routing": {"fallback_priority": priority}, "qa": {"browser_routing": True}}, + ) + return ModelConfig.objects.create( + provider=provider, + name=model_name, + display_name=f"QA Image {priority}", + capability=ModelConfig.Capability.IMAGE, + endpoint="images/generations", + unit_price=Decimal("20"), + status=ModelConfig.Status.ACTIVE, + metadata=_image_metadata(priority=priority), + ) + + +def _create_text_model(provider_name: str, priority: int, model_name: str) -> ModelConfig: + provider = ModelProvider.objects.create( + name=provider_name, + display_name=f"QA Provider {priority}", + status=ModelProvider.Status.ACTIVE, + base_url="http://127.0.0.1.invalid/v1", + api_key="", + metadata={"routing": {"fallback_priority": priority}, "qa": {"browser_routing": True}}, + ) + return ModelConfig.objects.create( + provider=provider, + name=model_name, + display_name=f"QA Script {priority}", + 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": ["streaming", "structured_output"], + }, + "pricing": {"base_cost_yuan": "0.10"}, + "qa": {"browser_routing_priority": priority}, + }, + ) + + +def _create_audio_model(provider_name: str, priority: int, model_name: str) -> ModelConfig: + provider = ModelProvider.objects.create( + name=provider_name, + display_name=f"QA Provider {priority}", + status=ModelProvider.Status.ACTIVE, + base_url="http://127.0.0.1.invalid/v1", + api_key="", + metadata={"routing": {"fallback_priority": priority}, "qa": {"browser_routing": True}}, + ) + return ModelConfig.objects.create( + provider=provider, + name=model_name, + display_name=f"QA Voice {priority}", + capability=ModelConfig.Capability.AUDIO, + endpoint="audio/speech", + unit_price=Decimal("10"), + status=ModelConfig.Status.ACTIVE, + metadata={ + "routing": {"fallback_on_failure": True, "fallback_candidate": True}, + "capabilities": { + "operations": ["tts"], + "features": [], + "languages": ["zh-CN"], + "voice_map": {"BV700_streaming": "alloy"}, + "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": "0.10", + }, + "qa": {"browser_routing_priority": priority}, + }, + ) + + +def _create_video_model(provider_name: str, priority: int, model_name: str) -> ModelConfig: + provider = ModelProvider.objects.create( + name=provider_name, + display_name=f"QA Provider {priority}", + status=ModelProvider.Status.ACTIVE, + base_url="http://127.0.0.1.invalid/v1", + api_key="", + metadata={"routing": {"fallback_priority": priority}, "qa": {"browser_routing": True}}, + ) + return ModelConfig.objects.create( + provider=provider, + name=model_name, + display_name=f"QA Video {priority}", + capability=ModelConfig.Capability.VIDEO, + endpoint="videos", + unit_price=Decimal("1"), + status=ModelConfig.Status.ACTIVE, + metadata={ + "routing": {"fallback_on_failure": True, "fallback_candidate": True}, + "capabilities": { + "operations": ["video_generate"], + "features": ["generate_audio"], + "max_reference_images": 9, + "max_reference_videos": 3, + "max_reference_audios": 3, + "aspect_ratios": ["9:16", "16:9"], + "resolutions": ["480p", "720p"], + "durations": [4, 15], + }, + "pricing": { + "unit": "cny_per_million_tokens", + "default": {"no_ref_video": 23, "with_ref_video": 14}, + }, + "qa": {"browser_routing_priority": priority}, + }, + ) + + +def _ledger_counts(task: AITask) -> dict[str, int]: + return { + kind: CreditLedger.objects.filter(task=task, ledger_type=kind).count() + for kind in ( + CreditLedger.Type.RESERVE, + CreditLedger.Type.CHARGE, + CreditLedger.Type.RELEASE, + ) + } + + +def setup(*, with_references: bool = False, tryon: bool = False, platform: bool = False) -> dict: + cleanup() + + user = User.objects.create(username=QA_USERNAME, status=User.Status.ACTIVE) + user.set_unusable_password() + user.save(update_fields=["password"]) + team = Team.objects.create(name=QA_TEAM_NAME, owner=user) + TeamMember.objects.create( + team=team, + user=user, + role=TeamMember.Role.OWNER, + status=TeamMember.Status.ACTIVE, + ) + CreditAccount.objects.create(team=team, balance=Decimal("1000"), reserved_balance=Decimal("0")) + + primary = _create_model(QA_PROVIDER_NAMES[0], -300, "qa-primary-image") + fallback_1 = _create_model(QA_PROVIDER_NAMES[1], -200, "qa-fallback-image-1") + fallback_2 = _create_model(QA_PROVIDER_NAMES[2], -100, "qa-fallback-image-2") + qa_model_ids = {primary.id, fallback_1.id, fallback_2.id} + + scenario = {"name": "first_success"} + provider_mocks: dict[object, Mock] = {} + + def provider_for(model: ModelConfig) -> Mock: + provider = provider_mocks.get(model.id) + if provider is None: + provider = Mock() + provider.extract_first_media_url.side_effect = lambda response: response["data"][0]["url"] + provider_mocks[model.id] = provider + + current = scenario["name"] + route_name = current.removeprefix("reference_").removeprefix("tryon_").removeprefix("platform_") + should_fail = route_name == "all_failed" or ( + route_name == "fallback_success" and model.id == primary.id + ) + if model.id not in qa_model_ids: + should_fail = True + provider.image_generation.side_effect = ( + requests.ConnectionError(f"QA controlled failure: {current}") + if should_fail + else None + ) + provider.image_generation.return_value = { + "data": [{"url": "http://qa.invalid/generated.png"}], + "usage": {"images": 1}, + } + provider.image_edit.side_effect = provider.image_generation.side_effect + provider.image_edit.return_value = provider.image_generation.return_value + return provider + + policy = deepcopy(settings.MODEL_ROUTING_POLICY) + policy["image"]["retry_delays"] = [0] + policy["jitter_ratio"] = 0 + + stored = StoredObject( + bucket="qa-browser", + object_key="qa/generated.png", + content_type="image/png", + size_bytes=8, + ) + + created: list[dict] = [] + with ( + override_settings(MODEL_ROUTING_POLICY=policy), + patch("apps.ai.services.get_default_model", return_value=primary), + patch("apps.ai.services.get_image_provider", side_effect=provider_for), + patch( + "apps.ai.services.VolcanoArkProvider.media_to_bytes", + return_value=(BytesIO(b"qa-image"), "image/png"), + ), + patch("apps.ai.services.TosStorage") as storage, + ): + storage.return_value.upload_fileobj.return_value = stored + reference_ids: list[str] = [] + tryon_product = None + tryon_model = None + if with_references or tryon or platform: + for index in range(2): + reference = Asset.objects.create( + team=team, + created_by=user, + name=f"QA 参考图 {index + 1}", + asset_type=Asset.Type.IMAGE, + source=Asset.Source.UPLOAD, + category=Asset.Category.UPLOAD, + ) + AssetFile.objects.create( + asset=reference, + object_key=f"qa/reference-{index + 1}.png", + bucket="qa-browser", + content_type="image/png", + preview_url=f"http://qa.invalid/reference-{index + 1}.png", + is_primary=True, + ) + reference_ids.append(str(reference.id)) + + if tryon: + product_asset = Asset.objects.get(id=reference_ids[0]) + product_asset.category = Asset.Category.PRODUCT_IMAGE + product_asset.save(update_fields=["category"]) + tryon_product = Product.objects.create( + team=team, + created_by=user, + title="QA 模特上身图商品", + category="服饰内衣", + cover_asset=product_asset, + ) + tryon_model = Asset.objects.get(id=reference_ids[1]) + tryon_model.source = Asset.Source.AI_GENERATED + tryon_model.category = Asset.Category.PERSON + tryon_model.save(update_fields=["source", "category"]) + + platform_product = None + if platform: + product_asset = Asset.objects.get(id=reference_ids[0]) + product_asset.category = Asset.Category.PRODUCT_IMAGE + product_asset.save(update_fields=["category"]) + platform_product = Product.objects.create( + team=team, + created_by=user, + title="QA 平台套图商品", + category="服饰内衣", + cover_asset=product_asset, + ) + + prefix = "platform_" if platform else ("tryon_" if tryon else ("reference_" if with_references else "")) + label = "平台套图" if platform else ("模特上身图" if tryon else ("多参考图" if with_references else "无参考图")) + for short_name, description in ( + ("first_success", "主模型首次成功"), + ("fallback_success", "主模型重试后 Fallback 成功"), + ("all_failed", "全部候选失败并释放预留"), + ): + name = f"{prefix}{short_name}" + prompt = f"QA 浏览器验收({label}):{description}" + scenario["name"] = name + provider_mocks.clear() + task = enqueue_standalone_images( + team=team, + user=user, + prompt=prompt, + mode="cover" if platform else ("model" if tryon else "image"), + count=1, + ratio="1:1", + image_model=f"{primary.provider.name}:{primary.name}", + reference_image_ids=reference_ids if with_references and not tryon else [], + product_id=( + str(platform_product.id) + if platform_product is not None + else (str(tryon_product.id) if tryon_product is not None else None) + ), + model_id=str(tryon_model.id) if tryon_model is not None else None, + platform_id="taobao" if platform else None, + dispatch=False, + )[0] + run_standalone_image_task(task_id=str(task.id)) + task.refresh_from_db() + created.append( + { + "scenario": name, + "task_id": str(task.id), + "status": task.status, + "attempts": task.model_attempts.count(), + "operations": list(task.model_attempts.values_list("operation", flat=True)), + "reference_images": list( + task.model_attempts.values_list("request_summary__reference_images", flat=True) + ), + "ledger": _ledger_counts(task), + } + ) + + account = CreditAccount.objects.get(team=team) + return { + "team": QA_TEAM_NAME, + "tasks": created, + "account": { + "balance": str(account.balance), + "reserved_balance": str(account.reserved_balance), + }, + } + + +def setup_base() -> dict: + """创建商品、人物、场景三类基础资产的受控成功/Fallback/全失败任务。""" + cleanup() + + user = User.objects.create(username=QA_USERNAME, status=User.Status.ACTIVE) + user.set_unusable_password() + user.save(update_fields=["password"]) + team = Team.objects.create(name=QA_TEAM_NAME, owner=user) + TeamMember.objects.create( + team=team, + user=user, + role=TeamMember.Role.OWNER, + status=TeamMember.Status.ACTIVE, + ) + CreditAccount.objects.create(team=team, balance=Decimal("1000"), reserved_balance=Decimal("0")) + + primary = _create_model(QA_PROVIDER_NAMES[0], -300, "qa-primary-image") + fallback_1 = _create_model(QA_PROVIDER_NAMES[1], -200, "qa-fallback-image-1") + fallback_2 = _create_model(QA_PROVIDER_NAMES[2], -100, "qa-fallback-image-2") + qa_model_ids = {primary.id, fallback_1.id, fallback_2.id} + + cover = Asset.objects.create( + team=team, + created_by=user, + name="QA 商品主图", + asset_type=Asset.Type.IMAGE, + source=Asset.Source.UPLOAD, + category=Asset.Category.PRODUCT_IMAGE, + ) + AssetFile.objects.create( + asset=cover, + object_key="qa/base-product.png", + bucket="qa-browser", + content_type="image/png", + preview_url="http://qa.invalid/base-product.png", + is_primary=True, + ) + product = Product.objects.create( + team=team, + created_by=user, + title="QA 基础资产商品", + category="服饰内衣", + cover_asset=cover, + ) + project = Project.objects.create( + team=team, + created_by=user, + product=product, + name="QA 基础资产项目", + ) + + scenario = {"name": ""} + provider_mocks: dict[object, Mock] = {} + + def provider_for(model: ModelConfig) -> Mock: + provider = provider_mocks.get(model.id) + if provider is None: + provider = Mock() + provider.extract_first_media_url.side_effect = lambda response: response["data"][0]["url"] + provider_mocks[model.id] = provider + current = scenario["name"] + should_fail = current.endswith("all_failed") or ( + current.endswith("fallback_success") and model.id == primary.id + ) + if model.id not in qa_model_ids: + should_fail = True + failure = requests.ConnectionError(f"QA controlled failure: {current}") if should_fail else None + response = {"data": [{"url": "http://qa.invalid/generated.png"}], "usage": {"images": 1}} + provider.image_generation.side_effect = failure + provider.image_generation.return_value = response + provider.image_edit.side_effect = failure + provider.image_edit.return_value = response + return provider + + def store_media(**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"], + ) + + policy = deepcopy(settings.MODEL_ROUTING_POLICY) + policy["image"]["retry_delays"] = [0] + policy["jitter_ratio"] = 0 + created: list[dict] = [] + with ( + override_settings(MODEL_ROUTING_POLICY=policy), + patch("apps.ai.services.get_default_model", return_value=primary), + patch("apps.ai.services.get_image_provider", side_effect=provider_for), + patch("apps.ai.services._store_generated_media", side_effect=store_media), + patch("apps.ai.tasks.generate_base_asset_task.delay"), + patch("apps.assets.review.submit_asset_for_review"), + ): + for kind in ( + BaseAssetGroup.Kind.PRODUCT, + BaseAssetGroup.Kind.PERSON, + BaseAssetGroup.Kind.SCENE, + ): + for short_name, description in ( + ("first_success", "主模型首次成功"), + ("fallback_success", "主模型重试后 Fallback 成功"), + ("all_failed", "全部候选失败并释放预留"), + ): + name = f"base_{kind}_{short_name}" + scenario["name"] = name + provider_mocks.clear() + task = generate_base_asset( + project=project, + user=user, + kind=kind, + prompt=f"QA 浏览器验收({kind} 基础资产):{description}", + label=f"QA-{kind}-{short_name}", + ) + run_base_asset_task(task_id=str(task.id)) + task.refresh_from_db() + created.append( + { + "scenario": name, + "task_id": str(task.id), + "status": task.status, + "attempts": task.model_attempts.count(), + "operations": list(task.model_attempts.values_list("operation", flat=True)), + "reference_images": list( + task.model_attempts.values_list( + "request_summary__reference_images", flat=True + ) + ), + "ledger": _ledger_counts(task), + } + ) + + account = CreditAccount.objects.get(team=team) + return { + "team": QA_TEAM_NAME, + "tasks": created, + "account": { + "balance": str(account.balance), + "reserved_balance": str(account.reserved_balance), + }, + } + + +def setup_model_triview(*, project_scope: bool = False) -> dict: + """创建团队模特或项目角色三视图的受控成功/Fallback/全失败任务。""" + cleanup() + user = User.objects.create(username=QA_USERNAME, status=User.Status.ACTIVE) + user.set_unusable_password() + user.save(update_fields=["password"]) + team = Team.objects.create(name=QA_TEAM_NAME, owner=user) + TeamMember.objects.create( + team=team, + user=user, + role=TeamMember.Role.OWNER, + status=TeamMember.Status.ACTIVE, + ) + CreditAccount.objects.create(team=team, balance=Decimal("1000"), reserved_balance=Decimal("0")) + primary = _create_model(QA_PROVIDER_NAMES[0], -300, "qa-primary-image") + fallback_1 = _create_model(QA_PROVIDER_NAMES[1], -200, "qa-fallback-image-1") + fallback_2 = _create_model(QA_PROVIDER_NAMES[2], -100, "qa-fallback-image-2") + qa_model_ids = {primary.id, fallback_1.id, fallback_2.id} + + portrait = Asset.objects.create( + team=team, + created_by=user, + name="QA 模特形象", + asset_type=Asset.Type.IMAGE, + source=Asset.Source.UPLOAD, + category=Asset.Category.MODEL_PORTRAIT, + in_library=False, + ) + AssetFile.objects.create( + asset=portrait, + object_key="qa/model-portrait.png", + bucket="qa-browser", + content_type="image/png", + preview_url="http://qa.invalid/model-portrait.png", + is_primary=True, + ) + model = None + project = None + if project_scope: + product = Product.objects.create( + team=team, + created_by=user, + title="QA 项目角色三视图商品", + ) + project = Project.objects.create( + team=team, + created_by=user, + product=product, + name="QA 项目角色三视图", + ) + else: + model = Model.objects.create( + team=team, + created_by=user, + name="QA 团队模特", + portrait_asset=portrait, + ) + scenario = {"name": ""} + provider_mocks: dict[object, Mock] = {} + + def provider_for(model_config: ModelConfig) -> Mock: + provider = provider_mocks.get(model_config.id) + if provider is None: + provider = Mock() + provider.extract_first_media_url.side_effect = lambda response: response["data"][0]["url"] + provider_mocks[model_config.id] = provider + current = scenario["name"] + should_fail = current.endswith("all_failed") or ( + current.endswith("fallback_success") and model_config.id == primary.id + ) + if model_config.id not in qa_model_ids: + should_fail = True + provider.image_edit.side_effect = ( + requests.ConnectionError(f"QA controlled failure: {current}") if should_fail else None + ) + provider.image_edit.return_value = { + "data": [{"url": "http://qa.invalid/model-triview.png"}], + "usage": {"images": 1}, + } + return provider + + policy = deepcopy(settings.MODEL_ROUTING_POLICY) + policy["image"]["retry_delays"] = [0] + policy["jitter_ratio"] = 0 + stored = StoredObject( + bucket="qa-browser", + object_key="qa/model-triview.png", + content_type="image/png", + size_bytes=8, + ) + + def store_media(**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"], + ) + + created: list[dict] = [] + with ( + override_settings(MODEL_ROUTING_POLICY=policy), + patch("apps.ai.services.get_default_model", return_value=primary), + patch("apps.ai.services.get_image_provider", side_effect=provider_for), + patch("apps.ai.tasks.generate_model_triview_task.delay"), + patch("apps.ai.tasks.generate_triview_task.delay"), + patch("apps.ai.services._store_generated_media", side_effect=store_media), + patch( + "apps.ai.services.VolcanoArkProvider.media_to_bytes", + return_value=(BytesIO(b"qa-image"), "image/png"), + ), + patch("apps.ai.services.TosStorage") as storage, + patch("apps.assets.review.submit_asset_for_review"), + ): + storage.return_value.upload_fileobj.return_value = stored + for short_name, description in ( + ("first_success", "主模型首次成功"), + ("fallback_success", "主模型重试后 Fallback 成功"), + ("all_failed", "全部候选失败并释放预留"), + ): + prefix = "project_triview" if project_scope else "model_triview" + name = f"{prefix}_{short_name}" + scenario["name"] = name + provider_mocks.clear() + if project_scope: + task = generate_person_triview( + project=project, + user=user, + portrait_asset=portrait, + ) + run_triview_task(task_id=str(task.id)) + else: + task, task_created = generate_model_triview(model=model, user=user) + if not task_created: + raise RuntimeError("QA 模特三视图任务被错误复用") + run_model_triview_task(task_id=str(task.id)) + task.refresh_from_db() + created.append( + { + "scenario": name, + "description": description, + "task_id": str(task.id), + "status": task.status, + "attempts": task.model_attempts.count(), + "operations": list(task.model_attempts.values_list("operation", flat=True)), + "reference_images": list( + task.model_attempts.values_list( + "request_summary__reference_images", flat=True + ) + ), + "aspect_ratios": list( + task.model_attempts.values_list( + "request_summary__aspect_ratio", flat=True + ) + ), + "ledger": _ledger_counts(task), + } + ) + + account = CreditAccount.objects.get(team=team) + return { + "team": QA_TEAM_NAME, + "tasks": created, + "account": { + "balance": str(account.balance), + "reserved_balance": str(account.reserved_balance), + }, + } + + +def setup_storyboard() -> dict: + """创建单个故事板分镜的受控成功/Fallback/全失败任务。""" + cleanup() + user = User.objects.create(username=QA_USERNAME, status=User.Status.ACTIVE) + user.set_unusable_password() + user.save(update_fields=["password"]) + team = Team.objects.create(name=QA_TEAM_NAME, owner=user) + TeamMember.objects.create( + team=team, + user=user, + role=TeamMember.Role.OWNER, + status=TeamMember.Status.ACTIVE, + ) + CreditAccount.objects.create(team=team, balance=Decimal("1000"), reserved_balance=Decimal("0")) + primary = _create_model(QA_PROVIDER_NAMES[0], -300, "qa-primary-image") + fallback_1 = _create_model(QA_PROVIDER_NAMES[1], -200, "qa-fallback-image-1") + fallback_2 = _create_model(QA_PROVIDER_NAMES[2], -100, "qa-fallback-image-2") + qa_model_ids = {primary.id, fallback_1.id, fallback_2.id} + + product = Product.objects.create( + team=team, + created_by=user, + title="QA 故事板商品", + ) + project = Project.objects.create( + team=team, + created_by=user, + product=product, + name="QA 故事板项目", + ) + script = ScriptVersion.objects.create( + project=project, + title="QA 故事板脚本", + content="结构化脚本", + is_adopted=True, + ) + segment = ScriptSegment.objects.create( + script_version=script, + sort_order=0, + visual_prompt="QA 人物展示商品", + ) + shot = StoryboardShot.objects.create( + project=project, + script_segment=segment, + sort_order=0, + status=StoryboardShot.Status.QUEUED, + ) + refs = [ + {"url": "http://qa.invalid/storyboard-person.png", "label": "女主", "type": "person"}, + {"url": "http://qa.invalid/storyboard-product.png", "label": "商品", "type": "product"}, + ] + scenario = {"name": ""} + provider_mocks: dict[object, Mock] = {} + + def provider_for(model_config: ModelConfig) -> Mock: + provider = provider_mocks.get(model_config.id) + if provider is None: + provider = Mock() + provider.extract_first_media_url.side_effect = lambda response: response["data"][0]["url"] + provider_mocks[model_config.id] = provider + current = scenario["name"] + should_fail = current.endswith("all_failed") or ( + current.endswith("fallback_success") and model_config.id == primary.id + ) + if model_config.id not in qa_model_ids: + should_fail = True + provider.image_edit.side_effect = ( + requests.ConnectionError(f"QA controlled failure: {current}") if should_fail else None + ) + provider.image_edit.return_value = { + "data": [{"url": "http://qa.invalid/storyboard.png"}], + "usage": {"images": 1}, + } + return provider + + def store_media(**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"], + ) + + policy = deepcopy(settings.MODEL_ROUTING_POLICY) + policy["image"]["retry_delays"] = [0] + policy["jitter_ratio"] = 0 + created: list[dict] = [] + with ( + override_settings(MODEL_ROUTING_POLICY=policy, STORYBOARD_MAX_PARALLEL=1), + patch("apps.ai.services.get_default_model", return_value=primary), + patch("apps.ai.services.get_image_provider", side_effect=provider_for), + patch("apps.ai.services._storyboard_reference_images", return_value=refs), + patch( + "apps.ai.services.build_storyboard_frame_prompt_refs", + return_value="QA 带编号参考图故事板提示词", + ), + patch("apps.ai.services._store_generated_media", side_effect=store_media), + patch("apps.assets.review.submit_asset_for_review"), + patch("apps.ai.services.notify_generation_failure"), + patch("threading.Thread"), + ): + for short_name, description in ( + ("first_success", "主模型首次成功"), + ("fallback_success", "主模型重试后 Fallback 成功"), + ("all_failed", "全部候选失败并释放预留"), + ): + name = f"storyboard_{short_name}" + scenario["name"] = name + provider_mocks.clear() + StoryboardShot.objects.filter(id=shot.id).update( + status=StoryboardShot.Status.QUEUED, + error_message="", + ) + poll_storyboard(project=project, user=user) + task = AITask.objects.filter( + project=project, + task_type=AITask.Type.STORYBOARD, + request_payload__storyboard_shot=str(shot.id), + ).latest("created_at") + _storyboard_shot_worker(str(task.id), str(shot.id), str(user.id)) + task.refresh_from_db() + shot.refresh_from_db() + created.append( + { + "scenario": name, + "description": description, + "task_id": str(task.id), + "status": task.status, + "attempts": task.model_attempts.count(), + "operations": list(task.model_attempts.values_list("operation", flat=True)), + "reference_images": list( + task.model_attempts.values_list( + "request_summary__reference_images", flat=True + ) + ), + "aspect_ratios": list( + task.model_attempts.values_list( + "request_summary__aspect_ratio", flat=True + ) + ), + "shot_versions": shot.versions.count(), + "ledger": _ledger_counts(task), + } + ) + + account = CreditAccount.objects.get(team=team) + return { + "team": QA_TEAM_NAME, + "tasks": created, + "account": { + "balance": str(account.balance), + "reserved_balance": str(account.reserved_balance), + }, + } + + +def setup_entity() -> dict: + """创建实体提取的受控成功/Fallback/全失败任务。""" + cleanup() + user = User.objects.create(username=QA_USERNAME, status=User.Status.ACTIVE) + user.set_unusable_password() + user.save(update_fields=["password"]) + team = Team.objects.create(name=QA_TEAM_NAME, owner=user) + TeamMember.objects.create( + team=team, + user=user, + role=TeamMember.Role.OWNER, + status=TeamMember.Status.ACTIVE, + ) + CreditAccount.objects.create(team=team, balance=Decimal("1000"), reserved_balance=Decimal("0")) + primary = _create_text_model(QA_PROVIDER_NAMES[0], -300, "qa-primary-script") + fallback_1 = _create_text_model(QA_PROVIDER_NAMES[1], -200, "qa-fallback-script-1") + fallback_2 = _create_text_model(QA_PROVIDER_NAMES[2], -100, "qa-fallback-script-2") + qa_model_ids = {primary.id, fallback_1.id, fallback_2.id} + + product = Product.objects.create(team=team, created_by=user, title="QA 实体提取商品") + project = Project.objects.create( + team=team, + created_by=user, + product=product, + name="QA 实体提取项目", + ) + script = ScriptVersion.objects.create( + project=project, + title="QA 实体提取脚本", + content="结构化脚本", + is_adopted=True, + ) + ScriptSegment.objects.create( + script_version=script, + sort_order=0, + narration="女主在客厅展示商品", + visual_prompt="女主站在客厅", + ) + valid_events = [ + { + "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"}, + ] + scenario = {"name": ""} + provider_mocks: dict[object, Mock] = {} + + def provider_for(model_config: ModelConfig) -> Mock: + provider = provider_mocks.get(model_config.id) + if provider is None: + provider = Mock() + provider_mocks[model_config.id] = provider + current = scenario["name"] + should_fail = current.endswith("all_failed") or ( + current.endswith("fallback_success") and model_config.id == primary.id + ) + if model_config.id not in qa_model_ids: + should_fail = True + provider.chat_completion_stream.side_effect = ( + requests.ConnectionError(f"QA controlled failure: {current}") if should_fail else None + ) + provider.chat_completion_stream.return_value = valid_events + return provider + + policy = deepcopy(settings.MODEL_ROUTING_POLICY) + policy["text"]["retry_delays"] = [0, 0] + policy["jitter_ratio"] = 0 + created: list[dict] = [] + with ( + override_settings(MODEL_ROUTING_POLICY=policy), + patch("apps.ai.services._resolve_extract_model_config", return_value=primary), + patch("apps.ai.services.get_text_provider", side_effect=provider_for), + patch("apps.ai.tasks.extract_entities_task.delay"), + ): + for short_name, description in ( + ("first_success", "主模型首次成功"), + ("fallback_success", "主模型重试后 Fallback 成功"), + ("all_failed", "全部候选失败并释放预留"), + ): + name = f"entity_{short_name}" + scenario["name"] = name + provider_mocks.clear() + task = submit_extract_entities(project=project, user=user) + run_extract_entities_task(task_id=str(task.id)) + task.refresh_from_db() + project.refresh_from_db() + created.append( + { + "scenario": name, + "description": description, + "task_id": str(task.id), + "status": task.status, + "attempts": task.model_attempts.count(), + "operations": list(task.model_attempts.values_list("operation", flat=True)), + "streaming": list( + task.model_attempts.values_list("request_summary__streaming", flat=True) + ), + "structured_output": list( + task.model_attempts.values_list( + "request_summary__structured_output", flat=True + ) + ), + "entities_extracted": bool((project.metadata or {}).get("entities_extracted")), + "ledger": _ledger_counts(task), + } + ) + + account = CreditAccount.objects.get(team=team) + return { + "team": QA_TEAM_NAME, + "tasks": created, + "account": { + "balance": str(account.balance), + "reserved_balance": str(account.reserved_balance), + }, + } + + +def setup_script() -> dict: + """创建单镜同步优化与整稿流式生成各自的成功/Fallback/全失败任务。""" + cleanup() + user = User.objects.create(username=QA_USERNAME, status=User.Status.ACTIVE) + user.set_unusable_password() + user.save(update_fields=["password"]) + team = Team.objects.create(name=QA_TEAM_NAME, owner=user) + TeamMember.objects.create( + team=team, + user=user, + role=TeamMember.Role.OWNER, + status=TeamMember.Status.ACTIVE, + ) + CreditAccount.objects.create(team=team, balance=Decimal("1000"), reserved_balance=Decimal("0")) + primary = _create_text_model(QA_PROVIDER_NAMES[0], -300, "qa-primary-script") + fallback_1 = _create_text_model(QA_PROVIDER_NAMES[1], -200, "qa-fallback-script-1") + fallback_2 = _create_text_model(QA_PROVIDER_NAMES[2], -100, "qa-fallback-script-2") + qa_model_ids = {primary.id, fallback_1.id, fallback_2.id} + + product = Product.objects.create(team=team, created_by=user, title="QA 脚本路由商品") + project = Project.objects.create( + team=team, + created_by=user, + product=product, + name="QA 脚本路由项目", + ) + base_draft = { + "hook": "QA 钩子", + "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=project, + title="QA 基准稿", + content=json.dumps(base_draft, ensure_ascii=False), + metadata={"hook": "QA 钩子", "entities": [], "total_duration": 30}, + ) + ScriptSegment.objects.create( + script_version=base, + sort_order=0, + duration_seconds=15, + narration="旧口播0", + visual_prompt="旧画面0", + ) + target = ScriptSegment.objects.create( + script_version=base, + sort_order=1, + duration_seconds=15, + narration="旧口播1", + visual_prompt="旧画面1", + ) + revised = deepcopy(base_draft) + revised["segments"][1]["narration"] = "QA 优化口播" + revised["segments"][1]["visual"] = "QA 优化画面" + raw = json.dumps(revised, ensure_ascii=False) + events = [ + {"type": "reasoning", "text": "QA 分析"}, + {"type": "delta", "text": raw}, + {"type": "done"}, + ] + scenario = {"name": ""} + provider_mocks: dict[object, Mock] = {} + + def provider_for(model_config: ModelConfig) -> Mock: + provider = provider_mocks.get(model_config.id) + if provider is None: + provider = Mock() + provider_mocks[model_config.id] = provider + current = scenario["name"] + should_fail = current.endswith("all_failed") or ( + current.endswith("fallback_success") and model_config.id == primary.id + ) + if model_config.id not in qa_model_ids: + should_fail = True + error = requests.ConnectionError(f"QA controlled failure: {current}") if should_fail else None + provider.chat_completion.side_effect = error + provider.chat_completion.return_value = {"usage": {"total_tokens": 42}} + provider.extract_text.return_value = raw + provider.chat_completion_stream.side_effect = error + provider.chat_completion_stream.return_value = events + return provider + + policy = deepcopy(settings.MODEL_ROUTING_POLICY) + policy["text"]["retry_delays"] = [0, 0] + policy["jitter_ratio"] = 0 + created: list[dict] = [] + with ( + override_settings(MODEL_ROUTING_POLICY=policy), + patch("apps.ai.services.get_text_provider", side_effect=provider_for), + ): + for entry, description in (("segment", "单镜同步优化"), ("stream", "整稿流式生成")): + for short_name, outcome in ( + ("first_success", "主模型首次成功"), + ("fallback_success", "主模型重试后 Fallback 成功"), + ("all_failed", "全部候选失败并释放预留"), + ): + name = f"script_{entry}_{short_name}" + scenario["name"] = name + provider_mocks.clear() + before_versions = ScriptVersion.objects.filter(project=project).count() + try: + if entry == "segment": + regenerate_segment_via_agent( + project=project, + user=user, + model_config=primary, + segment=target, + instruction="更有吸引力", + ) + else: + list( + stream_script_agent( + project=project, + user=user, + model_config=primary, + mode="auto", + user_prompt="生成一版 QA 脚本", + aspect_ratio="9:16", + total_duration=30, + ) + ) + except Exception: + pass + task_type = ( + AITask.Type.SCRIPT_OPTIMIZATION + if entry == "segment" + else AITask.Type.SCRIPT_GENERATION + ) + task = AITask.objects.filter(project=project, task_type=task_type).latest("created_at") + task.refresh_from_db() + created.append( + { + "scenario": name, + "description": f"{description}:{outcome}", + "task_id": str(task.id), + "status": task.status, + "attempts": task.model_attempts.count(), + "operations": list(task.model_attempts.values_list("operation", flat=True)), + "streaming": list( + task.model_attempts.values_list("request_summary__streaming", flat=True) + ), + "structured_output": list( + task.model_attempts.values_list( + "request_summary__structured_output", flat=True + ) + ), + "new_versions": ( + ScriptVersion.objects.filter(project=project).count() - before_versions + ), + "ledger": _ledger_counts(task), + } + ) + + account = CreditAccount.objects.get(team=team) + return { + "team": QA_TEAM_NAME, + "tasks": created, + "account": { + "balance": str(account.balance), + "reserved_balance": str(account.reserved_balance), + }, + } + + +def setup_voice() -> dict: + """创建当前直连锁定与未来 OpenAI 兼容候选的配音受控任务。""" + cleanup() + user = User.objects.create(username=QA_USERNAME, status=User.Status.ACTIVE) + user.set_unusable_password() + user.save(update_fields=["password"]) + team = Team.objects.create(name=QA_TEAM_NAME, owner=user) + TeamMember.objects.create( + team=team, + user=user, + role=TeamMember.Role.OWNER, + status=TeamMember.Status.ACTIVE, + ) + CreditAccount.objects.create(team=team, balance=Decimal("1000"), reserved_balance=Decimal("0")) + primary = _create_audio_model(QA_PROVIDER_NAMES[0], -300, "qa-primary-voice") + fallback_1 = _create_audio_model(QA_PROVIDER_NAMES[1], -200, "qa-fallback-voice-1") + fallback_2 = _create_audio_model(QA_PROVIDER_NAMES[2], -100, "qa-fallback-voice-2") + qa_model_ids = {primary.id, fallback_1.id, fallback_2.id} + product = Product.objects.create(team=team, created_by=user, title="QA 配音商品") + project = Project.objects.create(team=team, created_by=user, product=product, name="QA 配音项目") + scenario = {"name": ""} + provider_mocks: dict[object, Mock] = {} + + def provider_for(model_config: ModelConfig) -> Mock: + provider = provider_mocks.get(model_config.id) + if provider is None: + provider = Mock() + provider.configured = True + provider_mocks[model_config.id] = provider + current = scenario["name"] + should_fail = current.endswith("all_failed") or current.endswith("locked_failed") or ( + current.endswith("fallback_success") and model_config.id == primary.id + ) + if model_config.id not in qa_model_ids: + should_fail = True + provider.synthesize.side_effect = ( + requests.ConnectionError(f"QA controlled failure: {current}") if should_fail else None + ) + provider.synthesize.return_value = (b"qa-mp3", 1200) + return provider + + policy = deepcopy(settings.MODEL_ROUTING_POLICY) + policy["audio"]["retry_delays"] = [0] + policy["jitter_ratio"] = 0 + stored = Mock(object_key="qa-voice.mp3", bucket="qa", content_type="audio/mpeg", size_bytes=6) + created: list[dict] = [] + with ( + override_settings(MODEL_ROUTING_POLICY=policy), + patch("apps.ai.services.get_default_model", return_value=primary), + patch("apps.ai.services.get_audio_provider", side_effect=provider_for), + patch("apps.ai.services.TosStorage.upload_fileobj", return_value=stored), + ): + for name, description, outbound in ( + ("voice_locked_first_success", "当前直连首次成功", False), + ("voice_locked_failed", "当前直连失败只重试不外切", False), + ("voice_fallback_success", "未来打开开关后动态 Fallback 成功", True), + ("voice_all_failed", "未来候选全部失败并释放预留", True), + ): + routing = dict((primary.metadata or {}).get("routing") or {}) + routing["fallback_on_failure"] = outbound + primary.metadata = {**(primary.metadata or {}), "routing": routing} + primary.save(update_fields=["metadata", "updated_at"]) + scenario["name"] = name + provider_mocks.clear() + before_assets = Asset.objects.filter(team=team, asset_type=Asset.Type.AUDIO).count() + try: + synthesize_project_voiceover( + project=project, + user=user, + items=[{"index": 0, "text": "QA 配音旁白"}], + voice_type="BV700_streaming", + speed_ratio=1.0, + ) + except Exception: + pass + task = AITask.objects.filter(project=project, task_type=AITask.Type.VOICEOVER).latest("created_at") + task.refresh_from_db() + created.append( + { + "scenario": name, + "description": description, + "task_id": str(task.id), + "status": task.status, + "attempts": task.model_attempts.count(), + "models": list(task.model_attempts.values_list("model_name", flat=True)), + "operations": list(task.model_attempts.values_list("operation", flat=True)), + "new_assets": Asset.objects.filter(team=team, asset_type=Asset.Type.AUDIO).count() + - before_assets, + "ledger": _ledger_counts(task), + } + ) + + account = CreditAccount.objects.get(team=team) + return { + "team": QA_TEAM_NAME, + "tasks": created, + "account": { + "balance": str(account.balance), + "reserved_balance": str(account.reserved_balance), + }, + } + + +def setup_video() -> dict: + """创建视频片段提交、Fallback、状态未知与全部失败的受控任务。""" + cleanup() + user = User.objects.create(username=QA_USERNAME, status=User.Status.ACTIVE) + user.set_unusable_password() + user.save(update_fields=["password"]) + team = Team.objects.create(name=QA_TEAM_NAME, owner=user) + TeamMember.objects.create( + team=team, + user=user, + role=TeamMember.Role.OWNER, + status=TeamMember.Status.ACTIVE, + ) + CreditAccount.objects.create(team=team, balance=Decimal("10000"), reserved_balance=Decimal("0")) + primary = _create_video_model(QA_PROVIDER_NAMES[0], -300, "qa-primary-video") + fallback_1 = _create_video_model(QA_PROVIDER_NAMES[1], -200, "qa-fallback-video-1") + fallback_2 = _create_video_model(QA_PROVIDER_NAMES[2], -100, "qa-fallback-video-2") + qa_model_ids = {primary.id, fallback_1.id, fallback_2.id} + product = Product.objects.create(team=team, created_by=user, title="QA 视频商品") + project = Project.objects.create(team=team, created_by=user, product=product, name="QA 视频片段项目") + scenario = {"name": ""} + provider_mocks: dict[object, Mock] = {} + + def provider_for(model_config: ModelConfig) -> Mock: + provider = provider_mocks.setdefault(model_config.id, Mock()) + current = scenario["name"] + provider.create_video_task.side_effect = None + provider.create_video_task.return_value = { + "id": f"qa-{current}-{model_config.id}", + "status": "queued", + } + if current == "video_state_unknown" and model_config.id == primary.id: + provider.create_video_task.side_effect = requests.ReadTimeout("QA response lost") + elif current == "video_all_failed" or ( + current == "video_fallback_success" and model_config.id == primary.id + ) or model_config.id not in qa_model_ids: + provider.create_video_task.side_effect = requests.ConnectionError( + f"QA controlled failure: {current}" + ) + provider.poll_video_task.return_value = { + "status": "succeeded", + "usage": {"total_tokens": 300000}, + "content": {"video_url": "https://qa.invalid/video.mp4"}, + } + provider.extract_first_media_url.return_value = "https://qa.invalid/video.mp4" + return provider + + def store_media(**kwargs): + return Asset.objects.create( + team=kwargs["team"], + created_by=kwargs["user"], + name=kwargs["name"], + asset_type=Asset.Type.VIDEO, + source=Asset.Source.AI_GENERATED, + category=Asset.Category.VIDEO_CLIP, + origin_task=kwargs["task"], + ) + + policy = deepcopy(settings.MODEL_ROUTING_POLICY) + policy["video"]["submit_retry_delays"] = [0] + policy["jitter_ratio"] = 0 + created: list[dict] = [] + with ( + override_settings(MODEL_ROUTING_POLICY=policy), + patch("apps.ai.services.get_default_model", return_value=primary), + patch("apps.ai.services.get_video_provider", side_effect=provider_for), + patch("apps.ai.services._video_reference_images", return_value=[]), + patch("apps.ai.services.build_video_segment_prompt", return_value="QA 视频提示词"), + patch("apps.ai.services._store_generated_media", side_effect=store_media), + patch("apps.ai.services.notify_generation_failure"), + ): + for index, (name, description) in enumerate( + ( + ("video_first_success", "视频片段首次提交成功并成片"), + ("video_fallback_success", "视频片段原模型重试后动态 Fallback 并成片"), + ("video_state_unknown", "视频片段提交状态未知且禁止重提"), + ("video_all_failed", "视频片段所有候选明确失败并释放预留"), + ) + ): + scenario["name"] = name + provider_mocks.clear() + segment = VideoSegment.objects.create( + project=project, + sort_order=index, + target_duration_seconds=15, + ) + before_assets = Asset.objects.filter(team=team, category=Asset.Category.VIDEO_CLIP).count() + try: + submit_video_segment(video_segment=segment, user=user, prompt="QA") + poll_video_segment(video_segment=segment, user=user) + except Exception: + pass + task = AITask.objects.filter( + project=project, + task_type=AITask.Type.VIDEO_SEGMENT, + request_payload__video_segment_id=str(segment.id), + ).latest("created_at") + task.refresh_from_db() + created.append( + { + "scenario": name, + "description": description, + "task_id": str(task.id), + "status": task.status, + "attempts": task.model_attempts.count(), + "models": list(task.model_attempts.values_list("model_name", flat=True)), + "provider_task_id": task.provider_task_id, + "new_assets": Asset.objects.filter( + team=team, category=Asset.Category.VIDEO_CLIP + ).count() + - before_assets, + "ledger": _ledger_counts(task), + } + ) + + account = CreditAccount.objects.get(team=team) + return { + "team": QA_TEAM_NAME, + "tasks": created, + "account": {"balance": str(account.balance), "reserved_balance": str(account.reserved_balance)}, + } + + +def setup_free_video() -> dict: + """创建自由创作视频多素材提交与动态 Fallback 的受控任务。""" + from apps.ai.free_video import finalize_free_video, submit_free_video + + cleanup() + user = User.objects.create(username=QA_USERNAME, status=User.Status.ACTIVE) + user.set_unusable_password() + user.save(update_fields=["password"]) + team = Team.objects.create(name=QA_TEAM_NAME, owner=user) + TeamMember.objects.create( + team=team, + user=user, + role=TeamMember.Role.OWNER, + status=TeamMember.Status.ACTIVE, + ) + CreditAccount.objects.create(team=team, balance=Decimal("10000"), reserved_balance=Decimal("0")) + primary = _create_video_model(QA_PROVIDER_NAMES[0], -300, "qa-primary-free-video") + fallback_1 = _create_video_model(QA_PROVIDER_NAMES[1], -200, "qa-fallback-free-video-1") + fallback_2 = _create_video_model(QA_PROVIDER_NAMES[2], -100, "qa-fallback-free-video-2") + qa_model_ids = {primary.id, fallback_1.id, fallback_2.id} + scenario = {"name": ""} + provider_mocks: dict[object, Mock] = {} + + def provider_for(model_config: ModelConfig) -> Mock: + provider = provider_mocks.setdefault(model_config.id, Mock()) + current = scenario["name"] + provider.create_video_task.side_effect = None + provider.create_video_task.return_value = { + "id": f"qa-{current}-{model_config.id}", + "status": "queued", + } + if current == "free_video_state_unknown" and model_config.id == primary.id: + provider.create_video_task.side_effect = requests.ReadTimeout("QA response lost") + elif current == "free_video_all_failed" or ( + current == "free_video_fallback_success" and model_config.id == primary.id + ) or model_config.id not in qa_model_ids: + provider.create_video_task.side_effect = requests.ConnectionError( + f"QA controlled failure: {current}" + ) + provider.poll_video_task.return_value = { + "status": "succeeded", + "usage": {"total_tokens": 30000}, + "content": {"video_url": "https://qa.invalid/free-video.mp4"}, + } + provider.extract_first_media_url.return_value = "https://qa.invalid/free-video.mp4" + return provider + + def store_media(*, task, media): + del media + return Asset.objects.create( + team=task.team, + created_by=task.created_by, + name="QA 自由创作视频", + asset_type=Asset.Type.VIDEO, + source=Asset.Source.AI_GENERATED, + category=Asset.Category.FREE_CREATE, + origin_task=task, + ) + + references = [ + {"url": "https://qa.invalid/ref.png", "type": "image", "label": "图片"}, + {"url": "https://qa.invalid/ref.mp4", "type": "video", "label": "视频", "duration": 1}, + {"url": "https://qa.invalid/ref.mp3", "type": "audio", "label": "音频"}, + ] + policy = deepcopy(settings.MODEL_ROUTING_POLICY) + policy["video"]["submit_retry_delays"] = [0] + policy["jitter_ratio"] = 0 + created: list[dict] = [] + with ( + override_settings(MODEL_ROUTING_POLICY=policy), + patch("apps.ai.free_video.FREE_VIDEO_MODELS", {primary.name}), + patch("apps.ai.services.get_video_provider", side_effect=provider_for), + patch("apps.ai.tasks.poll_free_video_task.apply_async"), + patch("apps.ai.free_video._store_free_video_media", side_effect=store_media), + patch("apps.ai.free_video._notify_failure"), + ): + for name, description in ( + ("free_video_first_success", "自由创作多素材首次提交成功并成片"), + ("free_video_fallback_success", "自由创作多素材原模型重试后 Fallback 并成片"), + ("free_video_state_unknown", "自由创作提交状态未知且禁止重提"), + ("free_video_all_failed", "自由创作所有候选明确失败并释放预留"), + ): + scenario["name"] = name + provider_mocks.clear() + before_assets = Asset.objects.filter(team=team, category=Asset.Category.FREE_CREATE).count() + task = submit_free_video( + team=team, + user=user, + params={ + "prompt": "QA 多素材自由创作", + "mode": "universal", + "model": primary.name, + "aspect_ratio": "16:9", + "resolution": "480p", + "duration": 4, + "generate_audio": True, + "references": references, + }, + ) + if task.status in (AITask.Status.SUBMITTED, AITask.Status.POLLING): + task = finalize_free_video(task=task) + task.refresh_from_db() + created.append( + { + "scenario": name, + "description": description, + "task_id": str(task.id), + "status": task.status, + "attempts": task.model_attempts.count(), + "models": list(task.model_attempts.values_list("model_name", flat=True)), + "provider_task_id": task.provider_task_id, + "new_assets": Asset.objects.filter( + team=team, category=Asset.Category.FREE_CREATE + ).count() + - before_assets, + "ledger": _ledger_counts(task), + } + ) + + account = CreditAccount.objects.get(team=team) + return { + "team": QA_TEAM_NAME, + "tasks": created, + "account": {"balance": str(account.balance), "reserved_balance": str(account.reserved_balance)}, + } + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument( + "action", + choices=( + "setup", + "setup-reference", + "setup-tryon", + "setup-platform", + "setup-base", + "setup-model-triview", + "setup-project-triview", + "setup-storyboard", + "setup-entity", + "setup-script", + "setup-voice", + "setup-video", + "setup-free-video", + "cleanup", + ), + ) + args = parser.parse_args() + _assert_safe() + if args.action == "setup": + result = setup() + elif args.action == "setup-reference": + result = setup(with_references=True) + elif args.action == "setup-tryon": + result = setup(tryon=True) + elif args.action == "setup-platform": + result = setup(platform=True) + elif args.action == "setup-base": + result = setup_base() + elif args.action == "setup-model-triview": + result = setup_model_triview() + elif args.action == "setup-project-triview": + result = setup_model_triview(project_scope=True) + elif args.action == "setup-storyboard": + result = setup_storyboard() + elif args.action == "setup-entity": + result = setup_entity() + elif args.action == "setup-script": + result = setup_script() + elif args.action == "setup-voice": + result = setup_voice() + elif args.action == "setup-video": + result = setup_video() + elif args.action == "setup-free-video": + result = setup_free_video() + else: + result = {"cleaned": cleanup()} + print(json.dumps(result, ensure_ascii=False, indent=2)) + + +if __name__ == "__main__": + main() diff --git a/docs/todo/模型调用与动态Fallback-完成说明.md b/docs/todo/模型调用与动态Fallback-完成说明.md new file mode 100644 index 0000000..3eb74e5 --- /dev/null +++ b/docs/todo/模型调用与动态Fallback-完成说明.md @@ -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 私有协议供应商。