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.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):
r = self.ac.get(f"/api/admin/models/?provider={self.prov.id}")
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] = []
snapshots: list[dict] = []
resolved_library_assets: list[dict[str, str]] = []
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] = {}
image_n = video_n = audio_n = 0
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)
return "Image"
def _resolve_group_assets(group: FreeAssetGroup) -> list[tuple[str, str, float]]:
resolved: list[tuple[str, str, float]] = []
def _resolve_group_assets(group: FreeAssetGroup) -> list[FreeAsset]:
resolved: list[FreeAsset] = []
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):
continue # 未就绪的跳过
if fa.status != FreeAsset.Status.ACTIVE:
continue
resolved.append(
(f"asset://{_normalize_remote_asset_id(fa.remote_asset_id)}", fa.asset_type, fa.duration or 0.0)
)
resolved.append(fa)
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 []:
url = str(ref.get("url") or "")
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:
raise ValueError(f"素材「{label or fa.name}」尚未就绪,请稍后重试")
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")
if mode == "keyframe":
if kind != "image":
@@ -220,9 +232,16 @@ def build_content_items(*, team, prompt: str, mode: str, references: list) -> di
asset_list = group_cache[gid]
if not asset_list:
raise ValueError(f"素材「{label or '未命名'}」尚未就绪,请在素材库中确认状态为「可用」后重试")
for asset_url, asset_type, dur in asset_list:
kind = {"Video": "video", "Audio": "audio"}.get(asset_type, "image")
_push(kind, asset_url, "reference_video" if kind == "video" else ("reference_audio" if kind == "audio" else "reference_image"), dur)
for fa in asset_list:
kind = {"Video": "video", "Audio": "audio"}.get(fa.asset_type, "image")
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
# 直传素材(已上传 TOS 的直链)
@@ -260,6 +279,7 @@ def build_content_items(*, team, prompt: str, mode: str, references: list) -> di
"content_items": content_items,
"api_prompt": api_prompt,
"snapshots": snapshots,
"resolved_library_assets": resolved_library_assets,
"image_n": image_n,
"video_n": video_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:
"""僵尸回收(趁每次新提交顺手做,无需定时任务):
· RESERVED 超 10 分钟:没提交到火山就死(worker 崩溃/进程重启)→ 标失败退费;
@@ -494,6 +531,22 @@ def submit_free_video(*, team, user, params: dict) -> AITask:
provider_code=code,
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.error_code = (code or "CreateTaskError")[:64]
task.error_message = raw_message[:2000]
+172 -3
View File
@@ -8,13 +8,42 @@ from __future__ import annotations
from dataclasses import asdict, dataclass
import re
from typing import Any
from typing import Any, Mapping
import requests
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 = {
"script_generate": "脚本",
"image_generate": "图片",
@@ -62,6 +91,10 @@ class PublicGenerationError:
_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:
@@ -79,7 +112,7 @@ def _build_error(code: str, operation: str, *, reference_id: str | None = None)
"model_unavailable": ("contact_support", False, f"{label}暂时不可用", "当前生成能力暂不可用,请稍后再试。"),
"content_rejected": ("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, "可用积分不足", "充值后可继续生成。"),
"task_timeout": ("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]:
"""尽力抽取 HTTP 状态与供应商 code/message;失败时仅返回异常文本。"""
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)
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(
exc: Exception,
*,
operation: str,
provider_name: str = "",
provider_code: str = "",
provider_metadata: Mapping[str, Any] | None = None,
internal_kind: str = "",
reference_id: str | None = None,
) -> PublicGenerationError:
@@ -158,7 +300,11 @@ def classify_generation_error(
# 内容审核必须先于 HTTP 403 判断:部分供应商会以 403 返回 policy violation。
if _contains(signal, "moderation_blocked", "safety_violation", "safety system", "content_policy", "policyviolation", "sensitivecontentdetected"):
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)
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)
@@ -172,6 +318,15 @@ def classify_generation_error(
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"):
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"):
return _build_error("invalid_input", operation, reference_id=reference_id)
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"
}:
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(
RuntimeError(raw_message),
operation=operation or TASK_OPERATIONS.get(task_type, "image_generate"),
provider_code=str(getattr(task, "error_code", "") or ""),
provider_metadata=provider_metadata,
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 typing import Any
from apps.ai.generation_errors import provider_error_rule_errors
from apps.ai.models import ModelConfig, ModelProvider
from apps.ai.routing_policy import load_model_routing_policy
@@ -114,15 +115,15 @@ def provider_fallback_priority(provider: ModelProvider) -> int:
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 ()
errors = list(provider_error_rule_errors(metadata))
if "fallback_priority" in routing:
value = routing["fallback_priority"]
if isinstance(value, bool) or not isinstance(value, int) or not 0 <= value <= 1000:
errors.append("routing.fallback_priority 必须是 0 到 1000 的整数,数字越小越优先")
return tuple(errors)
def model_metadata_errors(capability: str, metadata: Any) -> tuple[str, ...]:
+75
View File
@@ -10,7 +10,9 @@ from django.test import TestCase
from rest_framework.test import APIClient
from apps.accounts.models import Team, TeamMember, User
from apps.adminpanel.serializers import AdminTaskDetailSerializer
from apps.assets.models import FreeAsset, FreeAssetGroup
from apps.assets.free_asset_state import REMOTE_ASSET_UNAVAILABLE_MESSAGE
from apps.ai.free_video import (
build_content_items,
finalize_free_video,
@@ -28,6 +30,7 @@ from apps.ai.video_pricing import (
get_token_price,
)
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
STANDARD = "doubao-seedance-2-0-260128"
@@ -250,6 +253,78 @@ class SubmitFreeVideoTests(TestCase):
reservation = CreditReservation.objects.get(task=task)
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):
with self.assertRaisesMessage(ValueError, "仅标准档"):
submit_free_video(
@@ -66,6 +66,133 @@ class GenerationErrorClassifierTests(SimpleTestCase):
)
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):
error = classify_generation_error(
_http_error(400, "", "gateway rejected request"), operation="image_generate"
@@ -69,3 +69,29 @@ class GenerationNotificationSafetyTests(TestCase):
data = NotificationSerializer(notification).data
self.assertEqual(data["body"], "脚本暂时不可用。")
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("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):
self.assertEqual(VOLCANO_PROVIDER["metadata"]["routing"]["fallback_priority"], 10)
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)
unread = serializers.SerializerMethodField()
project_name = serializers.CharField(source="project.name", read_only=True)
brief = serializers.SerializerMethodField()
body = serializers.SerializerMethodField()
metadata = serializers.SerializerMethodField()
@@ -48,8 +49,27 @@ class NotificationSerializer(serializers.ModelSerializer):
def get_unread(self, obj):
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):
"""兼容早期 AI 失败通知:不再向普通用户回传嵌入正文的上游原始异常。"""
current = self._current_generation_message(obj)
if current:
return current
body = obj.body or ""
marker = "—— 第三方服务商 API 返回的原始报错 ——"
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: "当前生成能力暂不可用,请稍后再试。" }),
content_rejected: (operation) => ({ title: `${operation}内容需要调整`, description: "内容未通过生成审核,请调整描述或素材后重试。" }),
invalid_input: (operation) => ({ title: `${operation}内容需检查`, description: "请检查描述、参数或素材格式后重试。" }),
asset_unavailable: () => ({ title: "素材暂不可用", description: "引用的素材不存在或暂不可用,请更换后重试。" }),
asset_unavailable: () => ({ title: "引用素材当前不可用", description: "请重新上传或更换素材后重试。" }),
user_credit_insufficient: () => ({ title: "可用积分不足", description: "充值后可继续生成。" }),
task_timeout: (operation) => ({ title: `${operation}生成时间过长`, description: "服务响应较慢,未完成生成,请重试。" }),
processing_failed: (operation) => ({ title: `${operation}结果处理失败`, description: "本次未生成可用结果,请重试。" }),