fix: 优化AI生成失败提示与素材失效处理

This commit is contained in:
hh
2026-07-21 14:02:10 +08:00
parent d0c3690d47
commit a99aae8231
12 changed files with 746 additions and 20 deletions
+37
View File
@@ -640,6 +640,43 @@ class AdminModelProviderTests(TestCase):
self.assertEqual(r.status_code, 200) self.assertEqual(r.status_code, 200)
self.assertEqual(r.data["status"], "disabled") self.assertEqual(r.data["status"], "disabled")
def test_provider_error_rules_are_validated_on_admin_write(self):
valid_metadata = {
"error_rules": [
{
"provider_code": "NewProviderCode",
"contains_all": ["reference token", "expired"],
"public_code": "asset_unavailable",
"enabled": True,
}
]
}
accepted = self.ac.patch(
f"/api/admin/providers/{self.prov.id}/",
{"metadata": valid_metadata},
format="json",
)
self.assertEqual(accepted.status_code, 200)
self.assertEqual(accepted.data["metadata"], valid_metadata)
rejected = self.ac.patch(
f"/api/admin/providers/{self.prov.id}/",
{
"metadata": {
"error_rules": [
{
"contains_all": [],
"public_code": "arbitrary_new_code",
}
]
}
},
format="json",
)
self.assertEqual(rejected.status_code, 400)
self.prov.refresh_from_db()
self.assertEqual(self.prov.metadata, valid_metadata)
def test_models_list_filter(self): def test_models_list_filter(self):
r = self.ac.get(f"/api/admin/models/?provider={self.prov.id}") r = self.ac.get(f"/api/admin/models/?provider={self.prov.id}")
self.assertEqual(r.status_code, 200) self.assertEqual(r.status_code, 200)
+62 -9
View File
@@ -109,8 +109,9 @@ def build_content_items(*, team, prompt: str, mode: str, references: list) -> di
""" """
content_items: list[dict] = [] content_items: list[dict] = []
snapshots: list[dict] = [] snapshots: list[dict] = []
resolved_library_assets: list[dict[str, str]] = []
seen_urls: set[str] = set() seen_urls: set[str] = set()
group_cache: dict[str, list[tuple[str, str, float]]] = {} group_cache: dict[str, list[FreeAsset]] = {}
label_to_placeholder: dict[str, str] = {} label_to_placeholder: dict[str, str] = {}
image_n = video_n = audio_n = 0 image_n = video_n = audio_n = 0
video_duration_total = 0.0 # 输入参考视频总时长(token 公式的输入项 + ≤15s 校验) video_duration_total = 0.0 # 输入参考视频总时长(token 公式的输入项 + ≤15s 校验)
@@ -147,18 +148,28 @@ def build_content_items(*, team, prompt: str, mode: str, references: list) -> di
content_items.append(item) content_items.append(item)
return "Image" return "Image"
def _resolve_group_assets(group: FreeAssetGroup) -> list[tuple[str, str, float]]: def _resolve_group_assets(group: FreeAssetGroup) -> list[FreeAsset]:
resolved: list[tuple[str, str, float]] = [] resolved: list[FreeAsset] = []
for fa in group.assets.exclude(remote_asset_id="").order_by("created_at"): for fa in group.assets.exclude(remote_asset_id="").order_by("created_at"):
if fa.status == FreeAsset.Status.PROCESSING and not _refresh_processing_free_asset(fa): if fa.status == FreeAsset.Status.PROCESSING and not _refresh_processing_free_asset(fa):
continue # 未就绪的跳过 continue # 未就绪的跳过
if fa.status != FreeAsset.Status.ACTIVE: if fa.status != FreeAsset.Status.ACTIVE:
continue continue
resolved.append( resolved.append(fa)
(f"asset://{_normalize_remote_asset_id(fa.remote_asset_id)}", fa.asset_type, fa.duration or 0.0)
)
return resolved return resolved
def _remember_library_asset(fa: FreeAsset) -> None:
if any(item["local_asset_id"] == str(fa.id) for item in resolved_library_assets):
return
resolved_library_assets.append(
{
"local_asset_id": str(fa.id),
"remote_asset_id": fa.remote_asset_id,
"submitted_remote_asset_id": _normalize_remote_asset_id(fa.remote_asset_id),
"asset_name": fa.name,
}
)
for ref in references or []: for ref in references or []:
url = str(ref.get("url") or "") url = str(ref.get("url") or "")
ref_type = str(ref.get("type") or "image") ref_type = str(ref.get("type") or "image")
@@ -199,6 +210,7 @@ def build_content_items(*, team, prompt: str, mode: str, references: list) -> di
if fa.status != FreeAsset.Status.ACTIVE or not fa.remote_asset_id: if fa.status != FreeAsset.Status.ACTIVE or not fa.remote_asset_id:
raise ValueError(f"素材「{label or fa.name}」尚未就绪,请稍后重试") raise ValueError(f"素材「{label or fa.name}」尚未就绪,请稍后重试")
resolved_url = f"asset://{_normalize_remote_asset_id(fa.remote_asset_id)}" resolved_url = f"asset://{_normalize_remote_asset_id(fa.remote_asset_id)}"
_remember_library_asset(fa)
kind = {"Video": "video", "Audio": "audio"}.get(fa.asset_type, "image") kind = {"Video": "video", "Audio": "audio"}.get(fa.asset_type, "image")
if mode == "keyframe": if mode == "keyframe":
if kind != "image": if kind != "image":
@@ -220,9 +232,16 @@ def build_content_items(*, team, prompt: str, mode: str, references: list) -> di
asset_list = group_cache[gid] asset_list = group_cache[gid]
if not asset_list: if not asset_list:
raise ValueError(f"素材「{label or '未命名'}」尚未就绪,请在素材库中确认状态为「可用」后重试") raise ValueError(f"素材「{label or '未命名'}」尚未就绪,请在素材库中确认状态为「可用」后重试")
for asset_url, asset_type, dur in asset_list: for fa in asset_list:
kind = {"Video": "video", "Audio": "audio"}.get(asset_type, "image") kind = {"Video": "video", "Audio": "audio"}.get(fa.asset_type, "image")
_push(kind, asset_url, "reference_video" if kind == "video" else ("reference_audio" if kind == "audio" else "reference_image"), dur) asset_url = f"asset://{_normalize_remote_asset_id(fa.remote_asset_id)}"
_push(
kind,
asset_url,
"reference_video" if kind == "video" else ("reference_audio" if kind == "audio" else "reference_image"),
fa.duration or 0.0,
)
_remember_library_asset(fa)
continue continue
# 直传素材(已上传 TOS 的直链) # 直传素材(已上传 TOS 的直链)
@@ -260,6 +279,7 @@ def build_content_items(*, team, prompt: str, mode: str, references: list) -> di
"content_items": content_items, "content_items": content_items,
"api_prompt": api_prompt, "api_prompt": api_prompt,
"snapshots": snapshots, "snapshots": snapshots,
"resolved_library_assets": resolved_library_assets,
"image_n": image_n, "image_n": image_n,
"video_n": video_n, "video_n": video_n,
"audio_n": audio_n, "audio_n": audio_n,
@@ -267,6 +287,23 @@ def build_content_items(*, team, prompt: str, mode: str, references: list) -> di
} }
def _unavailable_library_asset_target(resolved_assets: list[dict], raw_message: str) -> dict | None:
"""从本次已解析引用中精确定位 Provider 指出的失效素材;不按陌生字符串模糊查库。"""
message = str(raw_message or "").casefold()
matches = []
for item in resolved_assets:
submitted_id = str(item.get("submitted_remote_asset_id") or "").casefold()
original_id = str(item.get("remote_asset_id") or "").casefold()
if (submitted_id and submitted_id in message) or (original_id and original_id in message):
matches.append(item)
if len(matches) == 1:
return matches[0]
if not matches and len(resolved_assets) == 1:
return resolved_assets[0]
return None
def _reap_stale_free_video_tasks(*, team) -> None: def _reap_stale_free_video_tasks(*, team) -> None:
"""僵尸回收(趁每次新提交顺手做,无需定时任务): """僵尸回收(趁每次新提交顺手做,无需定时任务):
· RESERVED 超 10 分钟:没提交到火山就死(worker 崩溃/进程重启)→ 标失败退费; · RESERVED 超 10 分钟:没提交到火山就死(worker 崩溃/进程重启)→ 标失败退费;
@@ -494,6 +531,22 @@ def submit_free_video(*, team, user, params: dict) -> AITask:
provider_code=code, provider_code=code,
reference_id=str(task.id), reference_id=str(task.id),
) )
if public_error.code == "asset_unavailable":
target = _unavailable_library_asset_target(
built["resolved_library_assets"],
raw_message,
)
if target is not None:
try:
from apps.assets.free_asset_state import mark_remote_asset_unavailable
mark_remote_asset_unavailable(
team=team,
local_asset_id=target["local_asset_id"],
remote_asset_id=target["remote_asset_id"],
)
except Exception: # noqa: BLE001 - 素材状态同步是附带自愈,不能遮蔽任务失败与退费。
logger.exception("failed to mark unavailable free asset for task %s", task.id)
task.status = AITask.Status.FAILED task.status = AITask.Status.FAILED
task.error_code = (code or "CreateTaskError")[:64] task.error_code = (code or "CreateTaskError")[:64]
task.error_message = raw_message[:2000] task.error_message = raw_message[:2000]
+172 -3
View File
@@ -8,13 +8,42 @@ from __future__ import annotations
from dataclasses import asdict, dataclass from dataclasses import asdict, dataclass
import re import re
from typing import Any from typing import Any, Mapping
import requests import requests
DOMAIN = "generation" DOMAIN = "generation"
PUBLIC_ERROR_CODES = frozenset(
{
"provider_quota_exhausted",
"provider_rate_limited",
"provider_unavailable",
"provider_config_error",
"model_unavailable",
"content_rejected",
"invalid_input",
"asset_unavailable",
"user_credit_insufficient",
"task_timeout",
"processing_failed",
"unknown",
}
)
# Provider 配置只能把外部异常映射到已有公共分类;平台积分与本地后处理必须由业务边界显式标记。
PROVIDER_RULE_PUBLIC_CODES = PUBLIC_ERROR_CODES - {
"user_credit_insufficient",
"processing_failed",
}
PROVIDER_ERROR_RULE_FIELDS = frozenset(
{"provider_code", "contains_all", "public_code", "enabled"}
)
MAX_PROVIDER_ERROR_RULES = 50
MAX_PROVIDER_RULE_TERMS = 5
MAX_PROVIDER_RULE_TERM_LENGTH = 100
OPERATION_LABELS = { OPERATION_LABELS = {
"script_generate": "脚本", "script_generate": "脚本",
"image_generate": "图片", "image_generate": "图片",
@@ -62,6 +91,10 @@ class PublicGenerationError:
_VOLCANO_ERROR_RE = re.compile(r"火山报错\s*\[([^\]]*)\]\s*(.*)", re.S) _VOLCANO_ERROR_RE = re.compile(r"火山报错\s*\[([^\]]*)\]\s*(.*)", re.S)
_SPECIFIED_ASSET_NOT_FOUND_RE = re.compile(
r"\bspecified\s+asset\s+[^\s]+\s+is\s+not\s+found\b",
re.I,
)
def _operation_label(operation: str) -> str: def _operation_label(operation: str) -> str:
@@ -79,7 +112,7 @@ def _build_error(code: str, operation: str, *, reference_id: str | None = None)
"model_unavailable": ("contact_support", False, f"{label}暂时不可用", "当前生成能力暂不可用,请稍后再试。"), "model_unavailable": ("contact_support", False, f"{label}暂时不可用", "当前生成能力暂不可用,请稍后再试。"),
"content_rejected": ("revise_input", False, f"{label}内容需要调整", "内容未通过生成审核,请调整描述或素材后重试。"), "content_rejected": ("revise_input", False, f"{label}内容需要调整", "内容未通过生成审核,请调整描述或素材后重试。"),
"invalid_input": ("revise_input", False, f"{label}内容需检查", "请检查描述、参数或素材格式后重试。"), "invalid_input": ("revise_input", False, f"{label}内容需检查", "请检查描述、参数或素材格式后重试。"),
"asset_unavailable": ("revise_input", False, "素材暂不可用", "引用的素材不存在或暂不可用,请更换后重试。"), "asset_unavailable": ("revise_input", False, "引用素材当前不可用", "请重新上传或更换素材后重试。"),
"user_credit_insufficient": ("recharge", False, "可用积分不足", "充值后可继续生成。"), "user_credit_insufficient": ("recharge", False, "可用积分不足", "充值后可继续生成。"),
"task_timeout": ("retry", True, f"{label}生成时间过长", "服务响应较慢,未完成生成,请重试。"), "task_timeout": ("retry", True, f"{label}生成时间过长", "服务响应较慢,未完成生成,请重试。"),
"processing_failed": ("retry", True, f"{label}结果处理失败", "本次未生成可用结果,请重试。"), "processing_failed": ("retry", True, f"{label}结果处理失败", "本次未生成可用结果,请重试。"),
@@ -96,6 +129,21 @@ def _build_error(code: str, operation: str, *, reference_id: str | None = None)
) )
def public_error_for_code(
code: str,
*,
operation: str,
reference_id: str | None = None,
) -> PublicGenerationError:
"""按稳定公共错误码生成当前安全文案,供历史结构化记录重新投影。"""
return _build_error(
code if code in PUBLIC_ERROR_CODES else "unknown",
operation,
reference_id=reference_id,
)
def _response_details(exc: Exception) -> tuple[int | None, str, str]: def _response_details(exc: Exception) -> tuple[int | None, str, str]:
"""尽力抽取 HTTP 状态与供应商 code/message;失败时仅返回异常文本。""" """尽力抽取 HTTP 状态与供应商 code/message;失败时仅返回异常文本。"""
response = getattr(exc, "response", None) response = getattr(exc, "response", None)
@@ -124,12 +172,106 @@ def _contains(text: str, *tokens: str) -> bool:
return any(token in text for token in tokens) return any(token in text for token in tokens)
def provider_error_rule_errors(metadata: Any) -> tuple[str, ...]:
"""校验 ``ModelProvider.metadata.error_rules``;运行时遇到非法配置一律忽略。"""
root = metadata if isinstance(metadata, Mapping) else {}
rules = root.get("error_rules")
if rules is None:
return ()
if not isinstance(rules, list):
return ("error_rules 必须是数组",)
if len(rules) > MAX_PROVIDER_ERROR_RULES:
return (f"error_rules 最多配置 {MAX_PROVIDER_ERROR_RULES}",)
errors: list[str] = []
for index, rule in enumerate(rules):
prefix = f"error_rules[{index}]"
if not isinstance(rule, Mapping):
errors.append(f"{prefix} 必须是对象")
continue
unsupported = sorted(set(rule) - PROVIDER_ERROR_RULE_FIELDS)
if unsupported:
errors.append(f"{prefix} 不支持字段:{', '.join(unsupported)}")
enabled = rule.get("enabled", True)
if not isinstance(enabled, bool):
errors.append(f"{prefix}.enabled 必须是 true 或 false")
public_code = rule.get("public_code")
if public_code not in PROVIDER_RULE_PUBLIC_CODES:
errors.append(f"{prefix}.public_code 必须是已有的供应商公共错误码")
provider_code = rule.get("provider_code", "")
if not isinstance(provider_code, str):
errors.append(f"{prefix}.provider_code 必须是字符串")
provider_code = ""
provider_code = provider_code.strip()
if len(provider_code) > 128:
errors.append(f"{prefix}.provider_code 不能超过 128 个字符")
contains_all = rule.get("contains_all", [])
if not isinstance(contains_all, list):
errors.append(f"{prefix}.contains_all 必须是字符串数组")
contains_all = []
elif len(contains_all) > MAX_PROVIDER_RULE_TERMS:
errors.append(f"{prefix}.contains_all 最多 {MAX_PROVIDER_RULE_TERMS}")
normalized_terms: list[str] = []
for term_index, term in enumerate(contains_all):
if not isinstance(term, str) or not term.strip():
errors.append(f"{prefix}.contains_all[{term_index}] 必须是非空字符串")
continue
normalized = term.strip()
if len(normalized) > MAX_PROVIDER_RULE_TERM_LENGTH:
errors.append(
f"{prefix}.contains_all[{term_index}] 不能超过 {MAX_PROVIDER_RULE_TERM_LENGTH} 个字符"
)
normalized_terms.append(normalized)
if not provider_code and not normalized_terms:
errors.append(f"{prefix} 至少配置 provider_code 或 contains_all")
return tuple(errors)
def _configured_provider_error(
metadata: Any,
*,
provider_code: str,
signal: str,
operation: str,
reference_id: str | None,
) -> PublicGenerationError | None:
if provider_error_rule_errors(metadata):
return None
root = metadata if isinstance(metadata, Mapping) else {}
rules = root.get("error_rules")
if not isinstance(rules, list):
return None
normalized_code = provider_code.strip().casefold()
normalized_signal = signal.casefold()
for rule in rules:
if not isinstance(rule, Mapping) or rule.get("enabled", True) is not True:
continue
expected_code = str(rule.get("provider_code") or "").strip().casefold()
terms = [str(item).strip().casefold() for item in rule.get("contains_all", [])]
if expected_code and expected_code != normalized_code:
continue
if terms and not all(term in normalized_signal for term in terms):
continue
public_code = str(rule.get("public_code") or "")
if public_code in PROVIDER_RULE_PUBLIC_CODES:
return _build_error(public_code, operation, reference_id=reference_id)
return None
def classify_generation_error( def classify_generation_error(
exc: Exception, exc: Exception,
*, *,
operation: str, operation: str,
provider_name: str = "", provider_name: str = "",
provider_code: str = "", provider_code: str = "",
provider_metadata: Mapping[str, Any] | None = None,
internal_kind: str = "", internal_kind: str = "",
reference_id: str | None = None, reference_id: str | None = None,
) -> PublicGenerationError: ) -> PublicGenerationError:
@@ -158,7 +300,11 @@ def classify_generation_error(
# 内容审核必须先于 HTTP 403 判断:部分供应商会以 403 返回 policy violation。 # 内容审核必须先于 HTTP 403 判断:部分供应商会以 403 返回 policy violation。
if _contains(signal, "moderation_blocked", "safety_violation", "safety system", "content_policy", "policyviolation", "sensitivecontentdetected"): if _contains(signal, "moderation_blocked", "safety_violation", "safety system", "content_policy", "policyviolation", "sensitivecontentdetected"):
return _build_error("content_rejected", operation, reference_id=reference_id) return _build_error("content_rejected", operation, reference_id=reference_id)
if _contains(compact_code, "assetnotfound") or _contains(signal, "asset not found", "referenced asset not found"): if (
_contains(compact_code, "assetnotfound")
or _contains(signal, "asset not found", "referenced asset not found")
or _SPECIFIED_ASSET_NOT_FOUND_RE.search(signal)
):
return _build_error("asset_unavailable", operation, reference_id=reference_id) return _build_error("asset_unavailable", operation, reference_id=reference_id)
if _contains(compact_code, "modelnotfound", "modelnotavailable") or _contains(signal, "model not found", "model is not available", "does not support"): if _contains(compact_code, "modelnotfound", "modelnotavailable") or _contains(signal, "model not found", "model is not available", "does not support"):
return _build_error("model_unavailable", operation, reference_id=reference_id) return _build_error("model_unavailable", operation, reference_id=reference_id)
@@ -172,6 +318,15 @@ def classify_generation_error(
return _build_error("provider_unavailable", operation, reference_id=reference_id) return _build_error("provider_unavailable", operation, reference_id=reference_id)
if status == 401 or status == 403 or exc.__class__.__name__ == "TtsNotConfigured" or _contains(signal, "api_key", "unauthorized", "not configured", "credential"): if status == 401 or status == 403 or exc.__class__.__name__ == "TtsNotConfigured" or _contains(signal, "api_key", "unauthorized", "not configured", "credential"):
return _build_error("provider_config_error", operation, reference_id=reference_id) return _build_error("provider_config_error", operation, reference_id=reference_id)
configured = _configured_provider_error(
provider_metadata,
provider_code=provider_code,
signal=signal,
operation=operation,
reference_id=reference_id,
)
if configured is not None:
return configured
if _contains(compact_code, "invalidparameter", "invalidimage", "invalidvideo", "invalidaudio") or _contains(signal, "invalid_image", "invalid image", "invalid parameter", "image_file"): if _contains(compact_code, "invalidparameter", "invalidimage", "invalidvideo", "invalidaudio") or _contains(signal, "invalid_image", "invalid image", "invalid parameter", "image_file"):
return _build_error("invalid_input", operation, reference_id=reference_id) return _build_error("invalid_input", operation, reference_id=reference_id)
if _contains(compact_code, "postprocesserror"): if _contains(compact_code, "postprocesserror"):
@@ -189,9 +344,23 @@ def public_error_for_task(task, *, operation: str | None = None) -> PublicGenera
"product_image", "person_image", "scene_image" "product_image", "person_image", "scene_image"
}: }:
operation = "base_asset_generate" operation = "base_asset_generate"
provider_metadata: Mapping[str, Any] | None = None
attempts = getattr(task, "model_attempts", None)
if attempts is not None and hasattr(attempts, "filter"):
attempt = (
attempts.filter(status="failed")
.select_related("provider")
.order_by("-sequence")
.first()
)
provider_metadata = getattr(getattr(attempt, "provider", None), "metadata", None)
if provider_metadata is None:
model_config = getattr(task, "model_config", None)
provider_metadata = getattr(getattr(model_config, "provider", None), "metadata", None)
return classify_generation_error( return classify_generation_error(
RuntimeError(raw_message), RuntimeError(raw_message),
operation=operation or TASK_OPERATIONS.get(task_type, "image_generate"), operation=operation or TASK_OPERATIONS.get(task_type, "image_generate"),
provider_code=str(getattr(task, "error_code", "") or ""), provider_code=str(getattr(task, "error_code", "") or ""),
provider_metadata=provider_metadata,
reference_id=str(getattr(task, "id", "") or "") or None, reference_id=str(getattr(task, "id", "") or "") or None,
) )
+8 -7
View File
@@ -11,6 +11,7 @@ from dataclasses import dataclass, field
from datetime import datetime from datetime import datetime
from typing import Any from typing import Any
from apps.ai.generation_errors import provider_error_rule_errors
from apps.ai.models import ModelConfig, ModelProvider from apps.ai.models import ModelConfig, ModelProvider
from apps.ai.routing_policy import load_model_routing_policy from apps.ai.routing_policy import load_model_routing_policy
@@ -114,15 +115,15 @@ def provider_fallback_priority(provider: ModelProvider) -> int:
def provider_metadata_errors(metadata: Any) -> tuple[str, ...]: def provider_metadata_errors(metadata: Any) -> tuple[str, ...]:
"""校验供应商路由配置;供后台写入校验和运行前诊断共同复用。""" """校验供应商路由与错误展示配置;供后台写入校验和运行前诊断共同复用。"""
routing = _dict(_dict(metadata).get("routing")) routing = _dict(_dict(metadata).get("routing"))
if "fallback_priority" not in routing: errors = list(provider_error_rule_errors(metadata))
return () if "fallback_priority" in routing:
value = routing["fallback_priority"] value = routing["fallback_priority"]
if isinstance(value, bool) or not isinstance(value, int) or not 0 <= value <= 1000: if isinstance(value, bool) or not isinstance(value, int) or not 0 <= value <= 1000:
return ("routing.fallback_priority 必须是 0 到 1000 的整数,数字越小越优先",) errors.append("routing.fallback_priority 必须是 0 到 1000 的整数,数字越小越优先")
return () return tuple(errors)
def model_metadata_errors(capability: str, metadata: Any) -> tuple[str, ...]: def model_metadata_errors(capability: str, metadata: Any) -> tuple[str, ...]:
+75
View File
@@ -10,7 +10,9 @@ from django.test import TestCase
from rest_framework.test import APIClient from rest_framework.test import APIClient
from apps.accounts.models import Team, TeamMember, User from apps.accounts.models import Team, TeamMember, User
from apps.adminpanel.serializers import AdminTaskDetailSerializer
from apps.assets.models import FreeAsset, FreeAssetGroup from apps.assets.models import FreeAsset, FreeAssetGroup
from apps.assets.free_asset_state import REMOTE_ASSET_UNAVAILABLE_MESSAGE
from apps.ai.free_video import ( from apps.ai.free_video import (
build_content_items, build_content_items,
finalize_free_video, finalize_free_video,
@@ -28,6 +30,7 @@ from apps.ai.video_pricing import (
get_token_price, get_token_price,
) )
from apps.billing.models import CreditAccount, CreditLedger, CreditReservation from apps.billing.models import CreditAccount, CreditLedger, CreditReservation
from apps.ops.models import Notification
from apps.billing.pricing import quote_video_actual, quote_video_estimate, video_reserve_amount from apps.billing.pricing import quote_video_actual, quote_video_estimate, video_reserve_amount
STANDARD = "doubao-seedance-2-0-260128" STANDARD = "doubao-seedance-2-0-260128"
@@ -250,6 +253,78 @@ class SubmitFreeVideoTests(TestCase):
reservation = CreditReservation.objects.get(task=task) reservation = CreditReservation.objects.get(task=task)
self.assertEqual(reservation.status, CreditReservation.Status.RELEASED) self.assertEqual(reservation.status, CreditReservation.Status.RELEASED)
def test_missing_remote_library_asset_is_failed_once_without_retry_or_raw_user_leak(self):
group = FreeAssetGroup.objects.create(
team=self.team,
created_by=self.user,
name="南南88",
remote_group_id="group-missing",
)
asset = FreeAsset.objects.create(
group=group,
name="正面照.png",
url="http://x/front.png",
remote_asset_id="asset-20260709095511-mntv8",
asset_type=FreeAsset.Type.IMAGE,
status=FreeAsset.Status.ACTIVE,
)
raw_message = (
"The parameter `content[1].image_url.url` specified in the request is not valid: "
"The specified asset asset-20260709095511-mntv8 is not found. "
"Request id: 021784602722967387079dd26c99e87d331bd24b7eb1a4c698fe2"
)
self.provider.create_video_task.side_effect = RuntimeError(
f"火山报错 [InvalidParameter] {raw_message}"
)
task = submit_free_video(
team=self.team,
user=self.user,
params=self._params(
references=[
{
"url": asset.url,
"type": "image",
"label": asset.name,
"source": "library",
"asset_id": str(asset.id),
}
]
),
)
asset.refresh_from_db()
task.refresh_from_db()
public_task = serialize_free_video_task(task)
admin_task = AdminTaskDetailSerializer(task).data
notification = Notification.objects.get(dedupe_key=f"task:{task.id}:failed")
attempts = list(task.model_attempts.order_by("sequence"))
reservation = CreditReservation.objects.get(task=task)
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(task.error_code, "InvalidParameter")
self.assertEqual(task.error_message, raw_message)
self.assertEqual(task.actual_cost, Decimal("0"))
self.assertEqual(task.base_cost, Decimal("0"))
self.assertEqual(reservation.status, CreditReservation.Status.RELEASED)
self.assertEqual(len(attempts), 1)
self.assertFalse(attempts[0].is_retry)
self.assertFalse(attempts[0].is_fallback)
self.assertEqual(attempts[0].error_type, "asset_unavailable")
self.assertIn(asset.remote_asset_id, admin_task["error_message"])
self.assertIn(asset.remote_asset_id, admin_task["attempts"][0]["raw_error"])
self.assertEqual(admin_task["attempts"][0]["provider_name"], attempts[0].provider_name)
self.assertTrue(admin_task["attempts"][0]["provider_name"])
self.assertEqual(asset.status, FreeAsset.Status.FAILED)
self.assertEqual(asset.error_message, REMOTE_ASSET_UNAVAILABLE_MESSAGE)
self.assertEqual(public_task["error"]["code"], "asset_unavailable")
self.assertEqual(notification.metadata["generation_error"]["code"], "asset_unavailable")
self.assertEqual(notification.body, public_task["error_message"])
self.assertNotIn(asset.remote_asset_id, public_task["error_message"])
self.assertNotIn("Request id", public_task["error_message"])
self.assertNotIn(asset.remote_asset_id, notification.body)
self.assertNotIn("Request id", notification.body)
def test_fast_1080p_rejected(self): def test_fast_1080p_rejected(self):
with self.assertRaisesMessage(ValueError, "仅标准档"): with self.assertRaisesMessage(ValueError, "仅标准档"):
submit_free_video( submit_free_video(
@@ -66,6 +66,133 @@ class GenerationErrorClassifierTests(SimpleTestCase):
) )
self.assertEqual(error.code, expected) self.assertEqual(error.code, expected)
def test_asset_unavailable_variants_are_classified_before_invalid_parameter(self):
cases = [
("AssetNotFound", "asset not found"),
("InvalidParameter", "referenced asset not found"),
(
"InvalidParameter",
"The parameter `content[1].image_url.url` specified in the request is not valid: "
"The specified asset asset-20260709095511-mntv8 is not found. "
"Request id: 021784602722967387079dd26c99e87d331bd24b7eb1a4c698fe2",
),
]
for provider_code, message in cases:
with self.subTest(provider_code=provider_code, message=message):
error = classify_generation_error(
_http_error(400, provider_code, message),
operation="video_generate",
reference_id="asset-task",
)
self.assertEqual(error.code, "asset_unavailable")
self.assertEqual(error.action, "revise_input")
self.assertFalse(error.retryable)
self.assertEqual(error.reference_id, "asset-task")
self.assertEqual(
error.fallback_message,
"引用素材当前不可用:请重新上传或更换素材后重试。",
)
self.assertNotIn("asset-20260709095511-mntv8", error.fallback_message)
self.assertNotIn("Request id", error.fallback_message)
self.assertNotIn("InvalidParameter", error.fallback_message)
def test_plain_invalid_parameter_remains_invalid_input(self):
error = classify_generation_error(
_http_error(400, "InvalidParameter", "duration must be between 4 and 12"),
operation="video_generate",
)
self.assertEqual(error.code, "invalid_input")
def test_provider_metadata_rule_maps_new_signature_for_public_projection_only(self):
metadata = {
"error_rules": [
{
"provider_code": "VendorReferenceExpired",
"contains_all": ["material handle", "expired"],
"public_code": "asset_unavailable",
"enabled": True,
}
]
}
raw = _http_error(
400,
"VendorReferenceExpired",
"The material handle ref-123 expired before submission",
)
unconfigured = classify_generation_error(raw, operation="video_generate")
configured = classify_generation_error(
raw,
operation="video_generate",
provider_metadata=metadata,
reference_id="configured-task",
)
self.assertEqual(unconfigured.code, "unknown")
self.assertEqual(configured.code, "asset_unavailable")
self.assertEqual(configured.reference_id, "configured-task")
self.assertNotIn("ref-123", configured.fallback_message)
def test_provider_metadata_rule_cannot_override_trusted_content_rejection(self):
error = classify_generation_error(
_http_error(403, "PolicyViolation", "safety_violation"),
operation="video_generate",
provider_metadata={
"error_rules": [
{
"provider_code": "PolicyViolation",
"public_code": "invalid_input",
"enabled": True,
}
]
},
)
self.assertEqual(error.code, "content_rejected")
def test_disabled_provider_metadata_rule_is_ignored(self):
error = classify_generation_error(
_http_error(400, "VendorReferenceExpired", "material handle expired"),
operation="video_generate",
provider_metadata={
"error_rules": [
{
"provider_code": "VendorReferenceExpired",
"public_code": "asset_unavailable",
"enabled": False,
}
]
},
)
self.assertEqual(error.code, "unknown")
def test_task_projection_reads_rules_from_its_own_provider_only(self):
rule_metadata = {
"error_rules": [
{
"contains_all": ["reference token", "expired"],
"public_code": "asset_unavailable",
"enabled": True,
}
]
}
common = {
"id": "provider-rule-task",
"task_type": "free_video",
"error_code": "VendorParameter",
"error_message": "reference token ref-456 expired",
}
configured_task = SimpleNamespace(
**common,
model_config=SimpleNamespace(provider=SimpleNamespace(metadata=rule_metadata)),
)
other_provider_task = SimpleNamespace(
**common,
model_config=SimpleNamespace(provider=SimpleNamespace(metadata={})),
)
self.assertEqual(public_error_for_task(configured_task).code, "asset_unavailable")
self.assertEqual(public_error_for_task(other_provider_task).code, "unknown")
def test_unknown_bad_request_is_not_blamed_on_user(self): def test_unknown_bad_request_is_not_blamed_on_user(self):
error = classify_generation_error( error = classify_generation_error(
_http_error(400, "", "gateway rejected request"), operation="image_generate" _http_error(400, "", "gateway rejected request"), operation="image_generate"
@@ -69,3 +69,29 @@ class GenerationNotificationSafetyTests(TestCase):
data = NotificationSerializer(notification).data data = NotificationSerializer(notification).data
self.assertEqual(data["body"], "脚本暂时不可用。") self.assertEqual(data["body"], "脚本暂时不可用。")
self.assertNotIn("api_error", data["metadata"]) self.assertNotIn("api_error", data["metadata"])
def test_structured_historical_notification_uses_current_shared_copy(self):
notification = Notification.objects.create(
team=self.team,
recipient=self.user,
dedupe_key="historical-asset-copy",
notification_type=Notification.Type.TASK,
priority=Notification.Priority.ERR,
title="自由创作视频失败",
brief="素材暂不可用:引用的素材不存在或暂不可用,请更换后重试。",
body="素材暂不可用:引用的素材不存在或暂不可用,请更换后重试。",
metadata={
"generation_error": {
"domain": "generation",
"code": "asset_unavailable",
"operation": "video_generate",
"action": "revise_input",
"retryable": False,
}
},
)
data = NotificationSerializer(notification).data
expected = "引用素材当前不可用:请重新上传或更换素材后重试。"
self.assertEqual(data["brief"], expected)
self.assertEqual(data["body"], expected)
@@ -191,6 +191,40 @@ class ModelRequirementsMatchTests(SimpleTestCase):
self.assertTrue(any("resolutions" 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)) self.assertTrue(any("durations" in error for error in errors))
def test_provider_error_rule_metadata_is_strictly_validated(self):
valid = provider_metadata_errors(
{
"error_rules": [
{
"provider_code": "InvalidParameter",
"contains_all": ["specified asset", "not found"],
"public_code": "asset_unavailable",
"enabled": True,
}
]
}
)
invalid_shape = provider_metadata_errors({"error_rules": {"public_code": "asset_unavailable"}})
invalid_code = provider_metadata_errors(
{"error_rules": [{"contains_all": ["x"], "public_code": "new_dynamic_code"}]}
)
invalid_internal = provider_metadata_errors(
{"error_rules": [{"contains_all": ["balance"], "public_code": "user_credit_insufficient"}]}
)
invalid_empty_match = provider_metadata_errors(
{"error_rules": [{"provider_code": "", "contains_all": [], "public_code": "unknown"}]}
)
invalid_field = provider_metadata_errors(
{"error_rules": [{"contains_any": ["x"], "public_code": "unknown"}]}
)
self.assertEqual(valid, ())
self.assertTrue(any("必须是数组" in error for error in invalid_shape))
self.assertTrue(any("public_code" in error for error in invalid_code))
self.assertTrue(any("public_code" in error for error in invalid_internal))
self.assertTrue(any("至少配置" in error for error in invalid_empty_match))
self.assertTrue(any("不支持字段" in error for error in invalid_field))
def test_bootstrap_catalog_contains_persistent_routing_contracts(self): def test_bootstrap_catalog_contains_persistent_routing_contracts(self):
self.assertEqual(VOLCANO_PROVIDER["metadata"]["routing"]["fallback_priority"], 10) self.assertEqual(VOLCANO_PROVIDER["metadata"]["routing"]["fallback_priority"], 10)
self.assertEqual(YUNQI_PROVIDER["metadata"]["routing"]["fallback_priority"], 20) self.assertEqual(YUNQI_PROVIDER["metadata"]["routing"]["fallback_priority"], 20)
@@ -0,0 +1,82 @@
"""自由创作远端素材状态的幂等同步服务。"""
from __future__ import annotations
from dataclasses import dataclass
from uuid import UUID
from apps.accounts.models import Team
from .models import FreeAsset
REMOTE_ASSET_UNAVAILABLE_MESSAGE = "远端素材已失效,请重新上传或删除"
@dataclass(frozen=True, slots=True)
class MarkRemoteAssetUnavailableResult:
matched: bool
changed: bool
asset_id: str | None = None
asset_name: str = ""
def _valid_uuid(value) -> UUID | None:
if value in (None, ""):
return None
try:
return UUID(str(value))
except (TypeError, ValueError, AttributeError):
return None
def mark_remote_asset_unavailable(
*,
team: Team,
remote_asset_id: str | None = None,
local_asset_id: str | None = None,
reason_code: str = "asset_unavailable",
) -> MarkRemoteAssetUnavailableResult:
"""把当前团队内精确命中的单个素材标为失效;无法唯一定位时不写数据库。"""
if reason_code != "asset_unavailable":
return MarkRemoteAssetUnavailableResult(matched=False, changed=False)
remote_id = str(remote_asset_id or "").strip()
local_id = _valid_uuid(local_asset_id)
if local_asset_id not in (None, "") and local_id is None:
return MarkRemoteAssetUnavailableResult(matched=False, changed=False)
if local_id is None and not remote_id:
return MarkRemoteAssetUnavailableResult(matched=False, changed=False)
queryset = FreeAsset.objects.filter(group__team=team, group__is_deleted=False)
if local_id is not None:
queryset = queryset.filter(id=local_id)
if remote_id:
queryset = queryset.filter(remote_asset_id__iexact=remote_id)
matches = list(queryset.order_by("id")[:2])
if len(matches) != 1:
return MarkRemoteAssetUnavailableResult(matched=False, changed=False)
asset = matches[0]
changed = (
asset.status != FreeAsset.Status.FAILED
or asset.error_message != REMOTE_ASSET_UNAVAILABLE_MESSAGE
)
if changed:
asset.status = FreeAsset.Status.FAILED
asset.error_message = REMOTE_ASSET_UNAVAILABLE_MESSAGE
asset.save(update_fields=["status", "error_message", "updated_at"])
return MarkRemoteAssetUnavailableResult(
matched=True,
changed=changed,
asset_id=str(asset.id),
asset_name=asset.name,
)
__all__ = [
"MarkRemoteAssetUnavailableResult",
"REMOTE_ASSET_UNAVAILABLE_MESSAGE",
"mark_remote_asset_unavailable",
]
@@ -0,0 +1,102 @@
from django.test import TestCase
from apps.accounts.models import Team, User
from apps.assets.free_asset_state import (
REMOTE_ASSET_UNAVAILABLE_MESSAGE,
mark_remote_asset_unavailable,
)
from apps.assets.models import FreeAsset, FreeAssetGroup
class MarkRemoteAssetUnavailableTests(TestCase):
def setUp(self):
self.user = User.objects.create_user(username="asset-state", password="p")
self.team = Team.objects.create(name="Asset State", owner=self.user)
self.group = FreeAssetGroup.objects.create(
team=self.team,
created_by=self.user,
name="角色",
remote_group_id="group-state",
)
self.asset = FreeAsset.objects.create(
group=self.group,
name="正面照",
remote_asset_id="Asset-STATE-1",
asset_type=FreeAsset.Type.IMAGE,
status=FreeAsset.Status.ACTIVE,
)
def test_exact_match_marks_one_asset_failed_and_is_idempotent(self):
first = mark_remote_asset_unavailable(
team=self.team,
local_asset_id=str(self.asset.id),
remote_asset_id="asset-STATE-1",
)
second = mark_remote_asset_unavailable(
team=self.team,
local_asset_id=str(self.asset.id),
remote_asset_id="asset-STATE-1",
)
self.asset.refresh_from_db()
self.assertTrue(first.matched)
self.assertTrue(first.changed)
self.assertEqual(first.asset_name, "正面照")
self.assertTrue(second.matched)
self.assertFalse(second.changed)
self.assertEqual(self.asset.status, FreeAsset.Status.FAILED)
self.assertEqual(self.asset.error_message, REMOTE_ASSET_UNAVAILABLE_MESSAGE)
def test_remote_id_only_requires_a_unique_team_scoped_match(self):
result = mark_remote_asset_unavailable(team=self.team, remote_asset_id="asset-STATE-1")
self.assertTrue(result.matched)
duplicate = FreeAsset.objects.create(
group=self.group,
name="重复远端 ID",
remote_asset_id="asset-STATE-1",
asset_type=FreeAsset.Type.IMAGE,
status=FreeAsset.Status.ACTIVE,
)
duplicate_result = mark_remote_asset_unavailable(
team=self.team,
remote_asset_id="asset-STATE-1",
)
duplicate.refresh_from_db()
self.assertFalse(duplicate_result.matched)
self.assertEqual(duplicate.status, FreeAsset.Status.ACTIVE)
def test_cross_team_or_mismatched_identifiers_never_update(self):
other_user = User.objects.create_user(username="asset-other", password="p")
other_team = Team.objects.create(name="Other", owner=other_user)
cross_team = mark_remote_asset_unavailable(
team=other_team,
local_asset_id=str(self.asset.id),
remote_asset_id=self.asset.remote_asset_id,
)
mismatch = mark_remote_asset_unavailable(
team=self.team,
local_asset_id=str(self.asset.id),
remote_asset_id="asset-different",
)
invalid_id = mark_remote_asset_unavailable(
team=self.team,
local_asset_id="not-a-uuid",
)
self.asset.refresh_from_db()
self.assertFalse(cross_team.matched)
self.assertFalse(mismatch.matched)
self.assertFalse(invalid_id.matched)
self.assertEqual(self.asset.status, FreeAsset.Status.ACTIVE)
def test_non_asset_reason_is_ignored(self):
result = mark_remote_asset_unavailable(
team=self.team,
local_asset_id=str(self.asset.id),
reason_code="invalid_input",
)
self.asset.refresh_from_db()
self.assertFalse(result.matched)
self.assertEqual(self.asset.status, FreeAsset.Status.ACTIVE)
+20
View File
@@ -7,6 +7,7 @@ class NotificationSerializer(serializers.ModelSerializer):
type = serializers.CharField(source="notification_type", read_only=True) type = serializers.CharField(source="notification_type", read_only=True)
unread = serializers.SerializerMethodField() unread = serializers.SerializerMethodField()
project_name = serializers.CharField(source="project.name", read_only=True) project_name = serializers.CharField(source="project.name", read_only=True)
brief = serializers.SerializerMethodField()
body = serializers.SerializerMethodField() body = serializers.SerializerMethodField()
metadata = serializers.SerializerMethodField() metadata = serializers.SerializerMethodField()
@@ -48,8 +49,27 @@ class NotificationSerializer(serializers.ModelSerializer):
def get_unread(self, obj): def get_unread(self, obj):
return not obj.is_read return not obj.is_read
def _current_generation_message(self, obj):
metadata = obj.metadata if isinstance(obj.metadata, dict) else {}
error = metadata.get("generation_error")
if not isinstance(error, dict) or error.get("domain") != "generation":
return ""
from apps.ai.generation_errors import public_error_for_code
return public_error_for_code(
str(error.get("code") or "unknown"),
operation=str(error.get("operation") or "image_generate"),
reference_id=str(error.get("reference_id") or "") or None,
).fallback_message
def get_brief(self, obj):
return self._current_generation_message(obj) or obj.brief or ""
def get_body(self, obj): def get_body(self, obj):
"""兼容早期 AI 失败通知:不再向普通用户回传嵌入正文的上游原始异常。""" """兼容早期 AI 失败通知:不再向普通用户回传嵌入正文的上游原始异常。"""
current = self._current_generation_message(obj)
if current:
return current
body = obj.body or "" body = obj.body or ""
marker = "—— 第三方服务商 API 返回的原始报错 ——" marker = "—— 第三方服务商 API 返回的原始报错 ——"
return body.split(marker, 1)[0].strip() return body.split(marker, 1)[0].strip()
+1 -1
View File
@@ -69,7 +69,7 @@ const ERROR_COPY: Record<GenerationErrorCode, Copy> = {
model_unavailable: (operation) => ({ title: `${operation}暂时不可用`, description: "当前生成能力暂不可用,请稍后再试。" }), model_unavailable: (operation) => ({ title: `${operation}暂时不可用`, description: "当前生成能力暂不可用,请稍后再试。" }),
content_rejected: (operation) => ({ title: `${operation}内容需要调整`, description: "内容未通过生成审核,请调整描述或素材后重试。" }), content_rejected: (operation) => ({ title: `${operation}内容需要调整`, description: "内容未通过生成审核,请调整描述或素材后重试。" }),
invalid_input: (operation) => ({ title: `${operation}内容需检查`, description: "请检查描述、参数或素材格式后重试。" }), invalid_input: (operation) => ({ title: `${operation}内容需检查`, description: "请检查描述、参数或素材格式后重试。" }),
asset_unavailable: () => ({ title: "素材暂不可用", description: "引用的素材不存在或暂不可用,请更换后重试。" }), asset_unavailable: () => ({ title: "引用素材当前不可用", description: "请重新上传或更换素材后重试。" }),
user_credit_insufficient: () => ({ title: "可用积分不足", description: "充值后可继续生成。" }), user_credit_insufficient: () => ({ title: "可用积分不足", description: "充值后可继续生成。" }),
task_timeout: (operation) => ({ title: `${operation}生成时间过长`, description: "服务响应较慢,未完成生成,请重试。" }), task_timeout: (operation) => ({ title: `${operation}生成时间过长`, description: "服务响应较慢,未完成生成,请重试。" }),
processing_failed: (operation) => ({ title: `${operation}结果处理失败`, description: "本次未生成可用结果,请重试。" }), processing_failed: (operation) => ({ title: `${operation}结果处理失败`, description: "本次未生成可用结果,请重试。" }),