fix: 修正模型超时与结果未知重放

This commit is contained in:
hh
2026-07-21 15:34:29 +08:00
parent a99aae8231
commit f2114bc36e
22 changed files with 1221 additions and 107 deletions
-2
View File
@@ -93,8 +93,6 @@ MODEL_ROUTING_VIDEO_SUBMIT_TIMEOUT=120
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
-2
View File
@@ -85,8 +85,6 @@ MODEL_ROUTING_VIDEO_SUBMIT_TIMEOUT=120
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
-2
View File
@@ -326,8 +326,6 @@ MODEL_ROUTING_POLICY = {
"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),
},
+2 -31
View File
@@ -306,8 +306,8 @@ def _unavailable_library_asset_target(resolved_assets: list[dict], raw_message:
def _reap_stale_free_video_tasks(*, team) -> None:
"""僵尸回收(趁每次新提交顺手做,无需定时任务):
· RESERVED 超 10 分钟:没提交到火山就死(worker 崩溃/进程重启)→ 标失败退费;
· SUBMITTED/POLLING 超 2 小时:轮询链早已断且无人认领(正常出片 5-10 分钟)→ 标失败退费;
· RESERVED 超过视频提交阶段总时限:没提交到火山就死(worker 崩溃/进程重启)→ 标失败退费;
· SUBMITTED/POLLING 已有远端任务 ID,只能按供应商终态收尾,不按本地等待时间回收;
· POSTPROCESSING 超 30 分钟:转存/结算中途崩溃 → 标失败退费(火山可能已出片,平台承担该笔成本)。"""
now = timezone.now()
from .routing_policy import load_model_routing_policy
@@ -319,11 +319,6 @@ def _reap_stale_free_video_tasks(*, team) -> None:
{"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)},
@@ -669,30 +664,6 @@ def finalize_free_video(*, task: AITask) -> AITask:
from .services import get_video_provider
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")
+71
View File
@@ -90,6 +90,12 @@ class PublicGenerationError:
return payload
class ProviderOutcomeUnknownError(RuntimeError):
"""请求可能已被供应商处理,但本地没有拿到可确认结果。"""
outcome_unknown = True
_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",
@@ -168,6 +174,71 @@ def _response_details(exc: Exception) -> tuple[int | None, str, str]:
return status if isinstance(status, int) else None, code, message or raw
_AFTER_SEND_CONNECTION_SIGNALS = (
"connection aborted",
"connection reset",
"remote disconnected",
"remote end closed connection",
"broken pipe",
"incomplete read",
"response ended prematurely",
"unexpected eof",
"protocolerror",
)
def _exception_chain(exc: BaseException):
"""按 cause/context 遍历异常链,识别 requests 包装下的真实网络阶段。"""
seen: set[int] = set()
current: BaseException | None = exc
while current is not None and id(current) not in seen:
seen.add(id(current))
yield current
current = current.__cause__ or current.__context__
def is_provider_outcome_unknown(exc: Exception) -> bool:
"""判断失败是否发生在“供应商可能已处理、但本地无法确认结果”的阶段。
这里只判断能否安全重放,不改变面向用户的公共错误码。建立连接前的失败仍可按
现有策略重试;读取响应、响应体中断或网关 502/504 一律保守禁止自动重放。
"""
for item in _exception_chain(exc):
if getattr(item, "outcome_unknown", False) is True:
return True
response = getattr(item, "response", None)
status = getattr(response, "status_code", None)
if status in {502, 504}:
return True
if isinstance(item, requests.ConnectTimeout):
continue
if isinstance(item, requests.ReadTimeout):
return True
# requests.Timeout 基类无法证明发生在建立连接前,按结果不确定处理。
if isinstance(item, requests.Timeout):
return True
if isinstance(
item,
(
requests.exceptions.ChunkedEncodingError,
requests.exceptions.ContentDecodingError,
),
):
return True
invalid_json_error = getattr(requests.exceptions, "InvalidJSONError", None)
if invalid_json_error is not None and isinstance(item, invalid_json_error):
return True
if isinstance(item, requests.ConnectionError):
signal = f"{item.__class__.__name__} {item}".casefold()
if any(token in signal for token in _AFTER_SEND_CONNECTION_SIGNALS):
return True
return False
def _contains(text: str, *tokens: str) -> bool:
return any(token in text for token in tokens)
@@ -0,0 +1,39 @@
"""只读预览或恢复被旧成片硬超时误判失败的视频任务。"""
from __future__ import annotations
import json
from django.core.management.base import BaseCommand, CommandError
from apps.ai.video_timeout_recovery import reconcile_video_timeout_task
class Command(BaseCommand):
help = "核对旧 generation_timeout 误判的视频任务;默认只读,--apply 才恢复远端成功结果"
def add_arguments(self, parser):
parser.add_argument(
"--task-id",
action="append",
required=True,
help="要核对的 AITask UUID;可重复传入",
)
parser.add_argument(
"--apply",
action="store_true",
help="仅对远端已 succeeded 的指定任务执行幂等恢复;不会重新提交模型或追扣用户积分",
)
def handle(self, *args, **options):
failed = []
for task_id in options["task_id"]:
try:
result = reconcile_video_timeout_task(str(task_id), apply=bool(options["apply"]))
except Exception as exc: # noqa: BLE001 — 命令需继续报告同批其他显式任务
failed.append(f"{task_id}: {exc}")
self.stderr.write(self.style.ERROR(f"{task_id}: {exc}"))
continue
self.stdout.write(json.dumps(result, ensure_ascii=False, default=str, sort_keys=True))
if failed:
raise CommandError("".join(failed))
+38 -4
View File
@@ -13,7 +13,11 @@ 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.generation_errors import (
PublicGenerationError,
classify_generation_error,
is_provider_outcome_unknown,
)
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
@@ -46,6 +50,7 @@ class ErrorRoutingDecision:
retry_current: bool
fallback: bool
exclude_provider: bool = False
outcome_unknown: bool = False
_TERMINAL_ERRORS = {
@@ -65,9 +70,19 @@ _SECRET_PATTERNS = (
)
def decide_error_routing(error: PublicGenerationError) -> ErrorRoutingDecision:
def decide_error_routing(
error: PublicGenerationError,
*,
outcome_unknown: bool = False,
) -> ErrorRoutingDecision:
"""把统一安全错误映射成路由动作;业务入口不得复制错误字符串判断。"""
if outcome_unknown:
return ErrorRoutingDecision(
retry_current=False,
fallback=False,
outcome_unknown=True,
)
if error.code in _TERMINAL_ERRORS:
return ErrorRoutingDecision(retry_current=False, fallback=False)
if error.code in _FALLBACK_ONLY_ERRORS:
@@ -361,7 +376,11 @@ def execute_model_call(
except Exception as exc:
elapsed = monotonic() - call_started
public_error = classify(exc, current_model)
decision = decide_error_routing(public_error)
outcome_unknown = is_provider_outcome_unknown(exc)
decision = decide_error_routing(
public_error,
outcome_unknown=outcome_unknown,
)
try:
failed_metadata = _coerce_metadata(error_meta(exc, current_model))
except Exception as metadata_exc: # noqa: BLE001 - 审计摘要失败不能遮蔽原始 Provider 异常。
@@ -370,7 +389,22 @@ def execute_model_call(
)
# 异步视频若已经拿到 Provider 任务 ID,表示提交结果并非“明确未创建”;禁止盲目重提。
if requirements.capability == ModelConfig.Capability.VIDEO and failed_metadata.provider_task_id:
decision = ErrorRoutingDecision(retry_current=False, fallback=False)
outcome_unknown = True
decision = ErrorRoutingDecision(
retry_current=False,
fallback=False,
outcome_unknown=True,
)
if outcome_unknown:
failed_metadata = AttemptMetadata(
provider_task_id=failed_metadata.provider_task_id,
usage=failed_metadata.usage,
platform_cost=failed_metadata.platform_cost,
response_summary={
**failed_metadata.response_summary,
"outcome_unknown": True,
},
)
_finish_attempt(
attempt,
status=AIModelAttempt.Status.FAILED,
-9
View File
@@ -41,7 +41,6 @@ class VideoRoutingPolicy:
submit_timeout: float
submit_total_timeout: float
poll_request_timeout: float
generation_timeout: float
retry_after_cap: float
@@ -75,7 +74,6 @@ _LABELS = {
"submit_timeout": "视频单次提交超时",
"submit_total_timeout": "视频提交阶段总时限",
"poll_request_timeout": "视频单次轮询超时",
"generation_timeout": "视频成片等待总时限",
}
@@ -229,7 +227,6 @@ def _video_policy(raw: Any) -> VideoRoutingPolicy:
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,
@@ -241,17 +238,11 @@ def _video_policy(raw: Any) -> VideoRoutingPolicy:
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,
)
+8 -22
View File
@@ -18,7 +18,12 @@ from django.db import transaction
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.generation_errors import (
TASK_OPERATIONS,
ProviderOutcomeUnknownError,
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,
@@ -200,6 +205,7 @@ def execute_routed_image_request(
except Exception as exc:
# Provider 已返回并可能产生上游费用;响应解析失败仍需审计本次真实尝试成本。
candidate_quote = quote_flat(actual_model, team=task.team)
exc.outcome_unknown = True
exc.attempt_metadata = AttemptMetadata(
usage=actual_response.get("usage") if isinstance(actual_response, dict) else {},
platform_cost=candidate_quote.base_cost_yuan,
@@ -853,7 +859,7 @@ def execute_routed_audio_request(
)
class VideoSubmissionStateUnknown(RuntimeError):
class VideoSubmissionStateUnknown(ProviderOutcomeUnknownError):
"""视频提交可能已到达供应商但未拿到可靠任务 ID;禁止自动重提。"""
@@ -3326,26 +3332,6 @@ def poll_video_segment(*, video_segment: VideoSegment, user) -> VideoSegmentVers
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,
+3 -4
View File
@@ -63,7 +63,8 @@ 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),
总等待上限由统一视频 generation_timeout 配置控制finalize 幂等(POSTPROCESSING 认领),
最多 60 ( 30 分钟);用尽后只停止 Worker 兜底,不把业务任务判失败
finalize 幂等(POSTPROCESSING 认领),
与前端主动 poll 并存不双扣轮询本身出错不重试(max_retries=0),下一次自重排继续"""
from apps.ai.free_video import finalize_free_video
from apps.ai.models import AITask
@@ -82,9 +83,7 @@ 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
# 是否终止由 finalize_free_video 读取统一 generation_timeout 决定;这里不再维护一套
# “60 次轮询”硬编码,避免修改成片等待时限时 Worker 与 Web 轮询口径不一致。
if task.status in (AITask.Status.SUBMITTED, AITask.Status.POLLING):
if task.status in (AITask.Status.SUBMITTED, AITask.Status.POLLING) and attempt < 60:
poll_free_video_task.apply_async(args=[task_id, attempt + 1], countdown=30)
return task_id
@@ -132,6 +132,8 @@ class FreeVideoRoutingTests(TestCase):
]
task = submit_free_video(team=self.team, user=self.user, params=self.params(references=references))
task.submitted_at = timezone.now() - timedelta(seconds=1900)
task.save(update_fields=["submitted_at", "updated_at"])
self.assertEqual(task.status, AITask.Status.SUBMITTED)
self.assertEqual(task.provider_task_id, "free-fallback")
@@ -206,9 +208,31 @@ class FreeVideoRoutingTests(TestCase):
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(task.model_attempts.count(), 1)
self.assertTrue(task.model_attempts.get().response_summary["outcome_unknown"])
self.assertFalse(unused_provider.create_video_task.called)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
def test_submit_gateway_504_is_not_retried_or_fallbacked(self):
candidate = self.candidate("free-video-gateway-unused")
primary_provider = self._provider_for(self.primary)
unused_provider = self._provider_for(candidate)
response = requests.Response()
response.status_code = 504
primary_provider.create_video_task.side_effect = requests.HTTPError(
"gateway timeout",
response=response,
)
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.assertTrue(task.model_attempts.get().response_summary["outcome_unknown"])
self.assertFalse(unused_provider.create_video_task.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_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"}
@@ -225,9 +249,10 @@ class FreeVideoRoutingTests(TestCase):
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):
def test_long_running_generation_keeps_polling_without_release_or_resubmit(self):
provider = self._provider_for(self.primary)
provider.create_video_task.return_value = {"id": "free-slow", "status": "queued"}
provider.poll_video_task.return_value = {"status": "running"}
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"])
@@ -235,9 +260,14 @@ class FreeVideoRoutingTests(TestCase):
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(result.status, AITask.Status.POLLING)
provider.poll_video_task.assert_called_once_with(
endpoint=self.primary.endpoint,
provider_task_id="free-slow",
timeout=60.0,
)
self.assertEqual(provider.create_video_task.call_count, 1)
self.assertEqual(result.model_attempts.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"))
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
self.assertGreater(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
+34 -1
View File
@@ -2,7 +2,12 @@ import requests
from types import SimpleNamespace
from django.test import SimpleTestCase
from apps.ai.generation_errors import classify_generation_error, public_error_for_task
from apps.ai.generation_errors import (
ProviderOutcomeUnknownError,
classify_generation_error,
is_provider_outcome_unknown,
public_error_for_task,
)
from apps.ai.script_agent import _script_error_event
@@ -23,6 +28,34 @@ def _http_error(status, code="", message=""):
class GenerationErrorClassifierTests(SimpleTestCase):
def test_outcome_unknown_identifies_after_send_failures(self):
cases = [
requests.ReadTimeout("response lost"),
requests.Timeout("phase unknown"),
_http_error(502, message="bad gateway"),
_http_error(504, message="gateway timeout"),
requests.ConnectionError("Connection aborted: RemoteDisconnected"),
requests.exceptions.ChunkedEncodingError("response ended prematurely"),
ProviderOutcomeUnknownError("provider may have accepted request"),
]
for exc in cases:
with self.subTest(exc=exc):
self.assertTrue(is_provider_outcome_unknown(exc))
def test_outcome_unknown_keeps_pre_connect_and_explicit_failures_retryable(self):
cases = [
requests.ConnectTimeout("connect timed out"),
requests.ConnectionError("Failed to establish a new connection: connection refused"),
requests.ConnectionError("DNS lookup failed"),
_http_error(429, message="rate limited"),
_http_error(500, message="internal error"),
]
for exc in cases:
with self.subTest(exc=exc):
self.assertFalse(is_provider_outcome_unknown(exc))
def test_yunqi_payment_required_is_provider_quota_not_user_credit(self):
error = classify_generation_error(
_http_error(402, message="Payment Required"),
@@ -121,6 +121,81 @@ class RoutingExecutorTests(TestCase):
self.assertEqual([item.is_retry for item in attempts], [False, True])
self.assertEqual(attempts[1].previous_attempt_id, attempts[0].id)
def test_outcome_unknown_never_retries_or_resolves_fallback_candidates(self):
primary = self.model(self.provider("exec-unknown-primary", 100), "primary")
candidate = self.model(self.provider("exec-unknown-candidate", 10), "candidate")
cases = [
("read-timeout", requests.ReadTimeout("response lost")),
("generic-timeout", requests.Timeout("phase unknown")),
("after-send-reset", requests.ConnectionError("Connection aborted: RemoteDisconnected")),
]
for key, exc in cases:
with self.subTest(key=key):
task = self.task(primary, f"exec-unknown-{key}")
resolver = Mock(return_value=[candidate])
with self.assertRaises(type(exc)):
self.execute(
task=task,
primary=primary,
invoke=lambda model, timeout, error=exc: (_ for _ in ()).throw(error),
candidate_resolver=resolver,
error_metadata=lambda error, model: AttemptMetadata(
response_summary={"request_sent": True}
),
)
resolver.assert_not_called()
self.assertEqual(task.model_attempts.count(), 1)
attempt = task.model_attempts.get()
self.assertEqual(
attempt.response_summary,
{"request_sent": True, "outcome_unknown": True},
)
self.assertFalse(attempt.is_retry)
self.assertFalse(attempt.is_fallback)
def test_gateway_502_and_504_never_retry_or_fallback(self):
primary = self.model(self.provider("exec-gateway-primary", 100), "primary")
candidate = self.model(self.provider("exec-gateway-candidate", 10), "candidate")
for status in (502, 504):
with self.subTest(status=status):
task = self.task(primary, f"exec-gateway-{status}")
resolver = Mock(return_value=[candidate])
response = requests.Response()
response.status_code = status
error = requests.HTTPError(f"gateway {status}", response=response)
with self.assertRaises(requests.HTTPError):
self.execute(
task=task,
primary=primary,
invoke=lambda model, timeout, exc=error: (_ for _ in ()).throw(exc),
candidate_resolver=resolver,
)
resolver.assert_not_called()
self.assertEqual(task.model_attempts.count(), 1)
self.assertTrue(task.model_attempts.get().response_summary["outcome_unknown"])
def test_connect_timeout_remains_retryable(self):
primary = self.model(self.provider("exec-connect-timeout", 100), "primary")
task = self.task(primary, "exec-connect-timeout")
calls = []
def invoke(model, timeout):
calls.append(model.id)
if len(calls) == 1:
raise requests.ConnectTimeout("connect timed out")
return {"ok": True}
result = self.execute(task=task, primary=primary, invoke=invoke)
self.assertEqual(result.call_count, 2)
self.assertEqual(task.model_attempts.count(), 2)
self.assertNotIn("outcome_unknown", task.model_attempts.first().response_summary)
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")
@@ -257,6 +332,7 @@ class RoutingExecutorTests(TestCase):
resolver.assert_not_called()
self.assertEqual(task.model_attempts.count(), 1)
self.assertEqual(task.model_attempts.get().provider_task_id, "provider-task-123")
self.assertTrue(task.model_attempts.get().response_summary["outcome_unknown"])
def test_metadata_extraction_failure_does_not_repeat_successful_model_call(self):
primary = self.model(self.provider("exec-meta", 100), "primary")
@@ -364,3 +440,25 @@ class RoutingExecutorTests(TestCase):
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"))
def test_outcome_unknown_keeps_single_reserve_and_single_release(self):
primary = self.model(self.provider("exec-unknown-billing", 100), "primary")
task = self.task(primary, "exec-unknown-billing")
reservation = reserve_credit(team=self.team, user=self.user, task=task, amount=Decimal("10"))
with self.assertRaises(requests.ReadTimeout):
self.execute(
task=task,
primary=primary,
invoke=lambda model, timeout: (_ for _ in ()).throw(
requests.ReadTimeout("response lost")
),
)
release_credit(reservation=reservation, reason="provider outcome unknown")
self.assertEqual(task.model_attempts.count(), 1)
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)
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
+8 -11
View File
@@ -35,7 +35,7 @@ class ModelRoutingPolicyTests(SimpleTestCase):
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.assertFalse(hasattr(policy.video, "generation_timeout"))
self.assertEqual(policy.postprocess.retry_delays, (2.0, 5.0, 10.0))
with self.assertRaises(FrozenInstanceError):
@@ -57,7 +57,6 @@ class ModelRoutingPolicyTests(SimpleTestCase):
submit_timeout=1,
submit_total_timeout=1,
poll_request_timeout=1,
generation_timeout=1,
retry_after_cap=0,
)
raw["postprocess"]["retry_delays"] = []
@@ -68,7 +67,7 @@ class ModelRoutingPolicyTests(SimpleTestCase):
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.video.poll_request_timeout, 1.0)
self.assertEqual(policy.postprocess.retry_delays, ())
def test_rejects_max_calls_smaller_than_model_count_with_chinese_message(self):
@@ -113,16 +112,14 @@ class ModelRoutingPolicyTests(SimpleTestCase):
):
load_model_routing_policy(raw)
def test_rejects_video_poll_timeout_longer_than_generation_window(self):
def test_video_poll_timeout_is_not_coupled_to_a_generation_window(self):
raw = self.policy_dict()
raw["video"]["poll_request_timeout"] = 61
raw["video"]["generation_timeout"] = 60
raw["video"]["poll_request_timeout"] = 600
with self.assertRaisesRegex(
RoutingPolicyConfigurationError,
"视频单次轮询超时.*不得大于.*generation_timeout",
):
load_model_routing_policy(raw)
policy = load_model_routing_policy(raw)
self.assertEqual(policy.video.poll_request_timeout, 600.0)
self.assertFalse(hasattr(policy.video, "generation_timeout"))
def test_app_startup_fails_fast_for_invalid_policy(self):
raw = self.policy_dict()
@@ -201,6 +201,30 @@ class ScriptOptimizationRoutingTests(TestCase):
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
def test_read_timeout_stops_without_retry_or_fallback_and_releases_once(self):
primary = self.model(self.provider("script-opt-unknown-primary", 100), "script-opt-primary")
candidate = self.model(
self.provider("script-opt-unknown-candidate", 10),
"script-opt-candidate",
outbound=False,
)
failed = self.good_provider()
failed.chat_completion.side_effect = requests.ReadTimeout("response lost")
self.provider_mocks[primary.id] = failed
with self.assertRaises(requests.ReadTimeout):
self.optimize(primary)
task = AITask.objects.get(task_type=AITask.Type.SCRIPT_OPTIMIZATION)
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(task.model_attempts.count(), 1)
self.assertTrue(task.model_attempts.get().response_summary["outcome_unknown"])
self.assertEqual(failed.chat_completion.call_count, 1)
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_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(
@@ -189,6 +189,31 @@ class ScriptStreamRoutingTests(TransactionTestCase):
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
def test_read_timeout_stops_stream_without_retry_or_fallback(self):
primary = self.model(self.provider("script-stream-unknown-primary", 100), "script-stream-primary")
candidate = self.model(
self.provider("script-stream-unknown-candidate", 10),
"script-stream-candidate",
outbound=False,
)
failed = Mock()
failed.chat_completion_stream.side_effect = requests.ReadTimeout("response lost")
self.provider_mocks[primary.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.assertEqual(task.model_attempts.count(), 1)
self.assertTrue(task.model_attempts.get().response_summary["outcome_unknown"])
self.assertEqual(failed.chat_completion_stream.call_count, 1)
self.assertNotIn(candidate.id, self.provider_mocks)
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)
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(
@@ -265,6 +265,57 @@ class StandaloneSingleImageRoutingTests(TestCase):
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
self.assertEqual(Asset.objects.filter(origin_task=task).count(), 1)
def test_read_timeout_stops_without_retry_or_fallback_and_releases_once(self):
primary = self.model(self.provider("single-unknown-primary", 100), "primary")
candidate = self.model(
self.provider("single-unknown-candidate", 10),
"candidate",
outbound=False,
)
task = self.submit(primary)[0]
failed = self._new_provider_mock()
failed.image_generation.side_effect = requests.ReadTimeout("response lost")
self.provider_mocks[primary.id] = 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.assertTrue(task.model_attempts.get().response_summary["outcome_unknown"])
self.assertEqual(failed.image_generation.call_count, 1)
self.assertNotIn(candidate.id, self.provider_mocks)
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)
def test_missing_image_result_is_outcome_unknown_and_is_not_replayed(self):
primary = self.model(self.provider("single-missing-primary", 100), "primary")
candidate = self.model(
self.provider("single-missing-candidate", 10),
"candidate",
outbound=False,
)
task = self.submit(primary)[0]
missing = self._new_provider_mock()
missing.image_generation.return_value = {"data": []}
self.provider_mocks[primary.id] = missing
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)
attempt = task.model_attempts.get()
self.assertTrue(attempt.response_summary["media_missing"])
self.assertTrue(attempt.response_summary["outcome_unknown"])
self.assertEqual(missing.image_generation.call_count, 1)
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_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")
@@ -151,6 +151,8 @@ class VideoSegmentRoutingTests(TestCase):
fallback_provider.create_video_task.return_value = {"id": "remote-fallback", "status": "queued"}
task = self._submit()
task.submitted_at = timezone.now() - timedelta(seconds=1900)
task.save(update_fields=["submitted_at", "updated_at"])
store_media.return_value = self._asset(task)
fallback_provider.poll_video_task.return_value = {
"status": "succeeded",
@@ -217,6 +219,7 @@ class VideoSegmentRoutingTests(TestCase):
task = AITask.objects.get(task_type=AITask.Type.VIDEO_SEGMENT, team=self.team)
self.assertEqual(task.model_attempts.count(), 1)
self.assertTrue(task.model_attempts.get().response_summary["outcome_unknown"])
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)
@@ -234,9 +237,34 @@ class VideoSegmentRoutingTests(TestCase):
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.assertTrue(attempt.response_summary["outcome_unknown"])
self.assertFalse(unused_provider.create_video_task.called)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
def test_submit_gateway_502_is_state_unknown_and_is_never_retried(self):
primary = self.model(self.provider("video-gateway", 20), "video-gateway", outbound=True, default=True)
candidate = self.model(self.provider("video-gateway-unused", 30), "video-gateway-unused", outbound=False)
primary_provider = self._provider_for(primary)
unused_provider = self._provider_for(candidate)
response = requests.Response()
response.status_code = 502
primary_provider.create_video_task.side_effect = requests.HTTPError(
"bad gateway",
response=response,
)
with self.assertRaises(requests.HTTPError):
self._submit()
task = AITask.objects.get(task_type=AITask.Type.VIDEO_SEGMENT, team=self.team)
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(task.model_attempts.count(), 1)
self.assertTrue(task.model_attempts.get().response_summary["outcome_unknown"])
self.assertFalse(unused_provider.create_video_task.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_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)
@@ -256,10 +284,11 @@ class VideoSegmentRoutingTests(TestCase):
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):
def test_long_running_generation_keeps_polling_without_release_or_resubmit(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"}
provider.poll_video_task.return_value = {"status": "running"}
task = self._submit()
task.submitted_at = timezone.now() - timedelta(seconds=1900)
task.save(update_fields=["submitted_at", "updated_at"])
@@ -269,10 +298,15 @@ class VideoSegmentRoutingTests(TestCase):
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(task.status, AITask.Status.POLLING)
self.assertEqual(self.segment.status, VideoSegment.Status.RUNNING)
provider.poll_video_task.assert_called_once_with(
endpoint=primary.endpoint,
provider_task_id="remote-slow",
timeout=60.0,
)
self.assertEqual(provider.create_video_task.call_count, 1)
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), 1)
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
self.assertGreater(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
@@ -0,0 +1,282 @@
from decimal import Decimal
from datetime import timedelta
from io import StringIO
from unittest.mock import Mock, patch
from django.core.management import call_command
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.video_timeout_recovery import reconcile_video_timeout_task
from apps.assets.models import Asset
from apps.billing.models import CreditAccount, CreditLedger, CreditReservation
from apps.billing.services.ledger import release_credit, reserve_credit
from apps.products.models import Product
from apps.projects.models import Project, VideoSegment, VideoSegmentVersion
class VideoTimeoutRecoveryTests(TestCase):
def setUp(self):
self.user = User.objects.create_user(username="video-timeout-recovery", password="x")
self.team = Team.objects.create(name="Video Timeout Recovery", 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"))
self.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=self.product,
name="Recovery project",
)
self.segment = VideoSegment.objects.create(
project=self.project,
sort_order=0,
target_duration_seconds=15,
status=VideoSegment.Status.FAILED,
error_message="视频遇到问题",
)
self.provider_config = ModelProvider.objects.create(
name="recovery-provider",
display_name="Recovery Provider",
status=ModelProvider.Status.ACTIVE,
base_url="https://video.example/v1",
)
self.model = ModelConfig.objects.create(
provider=self.provider_config,
name="recovery-video",
display_name="Recovery Video",
capability=ModelConfig.Capability.VIDEO,
endpoint="video/tasks",
status=ModelConfig.Status.ACTIVE,
metadata={
"pricing": {
"unit": "cny_per_million_tokens",
"default": {"no_ref_video": 46, "with_ref_video": 46},
}
},
)
def _released_task(self, *, task_type=AITask.Type.VIDEO_SEGMENT):
payload = {
"model": self.model.name,
"resolution": "720p",
"price_multiplier": "1",
"prompt": "recovery prompt",
}
project = self.project if task_type == AITask.Type.VIDEO_SEGMENT else None
if project:
payload["video_segment_id"] = str(self.segment.id)
else:
payload.update({"feature": "free_video", "mode": "universal", "references": []})
task = AITask.objects.create(
team=self.team,
created_by=self.user,
project=project,
task_type=task_type,
status=AITask.Status.FAILED,
model_config=self.model,
idempotency_key=f"recovery:{task_type}:{AITask.objects.count()}",
provider_task_id=f"remote-{task_type}",
request_payload=payload,
estimated_cost=Decimal("100"),
base_cost=Decimal("4.6"),
error_message="视频生成超过配置的等待总时限",
)
reservation = reserve_credit(team=self.team, user=self.user, task=task, amount=Decimal("100"))
release_credit(reservation=reservation, reason=task.error_message)
return task
@staticmethod
def _video_asset(task, category=Asset.Category.VIDEO_CLIP):
return Asset.objects.create(
team=task.team,
created_by=task.created_by,
origin_task=task,
name="recovered video",
asset_type=Asset.Type.VIDEO,
source=Asset.Source.AI_GENERATED,
category=category,
)
def test_command_defaults_to_read_only_preview(self):
task = self._released_task()
provider = Mock()
provider.poll_video_task.return_value = {"status": "running"}
output = StringIO()
with patch("apps.ai.services.get_video_provider", return_value=provider):
call_command("reconcile_video_timeouts", task_id=[str(task.id)], stdout=output)
task.refresh_from_db()
self.segment.refresh_from_db()
self.assertIn('"action": "preview"', output.getvalue())
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(self.segment.status, VideoSegment.Status.FAILED)
self.assertEqual(task.credit_reservation.status, CreditReservation.Status.RELEASED)
def test_apply_skips_remote_running_without_mutation(self):
task = self._released_task()
provider = Mock()
provider.poll_video_task.return_value = {"status": "running"}
with patch("apps.ai.services.get_video_provider", return_value=provider):
result = reconcile_video_timeout_task(str(task.id), apply=True)
task.refresh_from_db()
self.assertEqual(result["action"], "still_running")
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(task.error_message, "视频生成超过配置的等待总时限")
self.assertEqual(VideoSegmentVersion.objects.filter(task=task).count(), 0)
def test_apply_recovers_project_result_without_recharging_and_is_idempotent(self):
task = self._released_task()
provider = Mock()
provider.poll_video_task.return_value = {
"status": "succeeded",
"usage": {"total_tokens": 100000},
"content": {"video_url": "https://video.example/result.mp4"},
}
provider.extract_first_media_url.return_value = "https://video.example/result.mp4"
with (
patch("apps.ai.services.get_video_provider", return_value=provider),
patch("apps.ai.services._store_generated_media") as store_media,
):
store_media.side_effect = lambda **_kwargs: self._video_asset(task)
result = reconcile_video_timeout_task(str(task.id), apply=True)
again = reconcile_video_timeout_task(str(task.id), apply=True)
task.refresh_from_db()
self.segment.refresh_from_db()
self.assertEqual(result["action"], "recovered")
self.assertEqual(again["action"], "already_recovered")
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
self.assertEqual(task.actual_cost, Decimal("0"))
self.assertEqual(task.credit_reservation.status, CreditReservation.Status.RELEASED)
self.assertEqual(
CreditLedger.objects.filter(task=task, ledger_type=CreditLedger.Type.CHARGE).count(),
0,
)
self.assertEqual(VideoSegmentVersion.objects.filter(task=task).count(), 1)
self.assertEqual(self.segment.status, VideoSegment.Status.SUCCEEDED)
self.assertIsNotNone(self.segment.adopted_version_id)
self.assertEqual(store_media.call_count, 1)
self.assertEqual(provider.poll_video_task.call_count, 1)
def test_apply_recovers_free_video_without_recharging_and_is_idempotent(self):
task = self._released_task(task_type=AITask.Type.FREE_VIDEO)
provider = Mock()
provider.poll_video_task.return_value = {
"status": "succeeded",
"usage": {"total_tokens": 100000},
"content": {"video_url": "https://video.example/free.mp4"},
}
provider.extract_first_media_url.return_value = "https://video.example/free.mp4"
with (
patch("apps.ai.services.get_video_provider", return_value=provider),
patch("apps.ai.free_video._store_free_video_media") as store_media,
):
store_media.side_effect = lambda **_kwargs: self._video_asset(
task,
category=Asset.Category.FREE_CREATE,
)
result = reconcile_video_timeout_task(str(task.id), apply=True)
again = reconcile_video_timeout_task(str(task.id), apply=True)
task.refresh_from_db()
self.assertEqual(result["action"], "recovered")
self.assertEqual(again["action"], "already_recovered")
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
self.assertEqual(task.actual_cost, Decimal("0"))
self.assertEqual(task.credit_reservation.status, CreditReservation.Status.RELEASED)
self.assertEqual(
CreditLedger.objects.filter(task=task, ledger_type=CreditLedger.Type.CHARGE).count(),
0,
)
self.assertEqual(task.generated_assets.count(), 1)
self.assertEqual(store_media.call_count, 1)
self.assertEqual(provider.poll_video_task.call_count, 1)
def test_same_segment_only_recovers_one_timeout_result(self):
first = self._released_task()
second = self._released_task()
provider = Mock()
provider.poll_video_task.return_value = {
"status": "succeeded",
"usage": {"total_tokens": 100000},
"content": {"video_url": "https://video.example/result.mp4"},
}
provider.extract_first_media_url.return_value = "https://video.example/result.mp4"
with (
patch("apps.ai.services.get_video_provider", return_value=provider),
patch("apps.ai.services._store_generated_media") as store_media,
):
store_media.side_effect = lambda **_kwargs: self._video_asset(first)
recovered = reconcile_video_timeout_task(str(first.id), apply=True)
store_media.reset_mock()
skipped = reconcile_video_timeout_task(str(second.id), apply=True)
first.refresh_from_db()
second.refresh_from_db()
self.assertEqual(recovered["action"], "recovered")
self.assertEqual(skipped["action"], "skipped_existing_recovery")
self.assertEqual(first.status, AITask.Status.SUCCEEDED)
self.assertEqual(second.status, AITask.Status.FAILED)
self.assertEqual(VideoSegmentVersion.objects.filter(video_segment=self.segment).count(), 1)
store_media.assert_not_called()
def test_stale_reaper_does_not_fail_submitted_remote_task_by_elapsed_time(self):
from apps.ai.free_video import _reap_stale_free_video_tasks
task = AITask.objects.create(
team=self.team,
created_by=self.user,
task_type=AITask.Type.FREE_VIDEO,
status=AITask.Status.SUBMITTED,
model_config=self.model,
idempotency_key="recovery:stale-free-video",
provider_task_id="remote-still-running",
request_payload={"feature": "free_video"},
estimated_cost=Decimal("100"),
)
reserve_credit(team=self.team, user=self.user, task=task, amount=Decimal("100"))
old = timezone.now() - timedelta(days=2)
AITask.objects.filter(id=task.id).update(submitted_at=old, updated_at=old)
_reap_stale_free_video_tasks(team=self.team)
task.refresh_from_db()
task.credit_reservation.refresh_from_db()
self.assertEqual(task.status, AITask.Status.SUBMITTED)
self.assertEqual(task.credit_reservation.status, CreditReservation.Status.ACTIVE)
self.assertGreater(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
@override_settings(CELERY_TASK_ALWAYS_EAGER=False)
def test_free_video_worker_stops_rescheduling_after_sixty_without_failing_task(self):
from apps.ai.tasks import poll_free_video_task
task = AITask.objects.create(
team=self.team,
created_by=self.user,
task_type=AITask.Type.FREE_VIDEO,
status=AITask.Status.SUBMITTED,
model_config=self.model,
idempotency_key="recovery:worker-limit",
provider_task_id="remote-worker-limit",
)
with (
patch("apps.ai.free_video.finalize_free_video", return_value=task),
patch.object(poll_free_video_task, "apply_async") as enqueue,
):
poll_free_video_task.run(str(task.id), 59)
enqueue.assert_called_once_with(args=[str(task.id), 60], countdown=30)
enqueue.reset_mock()
poll_free_video_task.run(str(task.id), 60)
enqueue.assert_not_called()
task.refresh_from_db()
self.assertEqual(task.status, AITask.Status.SUBMITTED)
@@ -188,6 +188,33 @@ class VoiceoverRoutingTests(TestCase):
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
def test_read_timeout_stops_without_retry_or_fallback_and_releases_once(self):
primary = self.model(self.direct_provider, "voice-unknown-primary", outbound=True)
candidate = self.model(
self.provider("voice-unknown-candidate", 20),
"voice-unknown-fallback",
outbound=False,
voice="alloy",
)
self.direct_mock.synthesize.side_effect = requests.ReadTimeout("response lost")
fallback = Mock()
fallback.synthesize.return_value = (b"must-not-run", 1000)
self.provider_mocks[candidate.id] = fallback
with self.assertRaises(requests.ReadTimeout):
self.synthesize([{"index": 0, "text": "结果未知旁白"}])
task = AITask.objects.get(task_type=AITask.Type.VOICEOVER)
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(task.model_attempts.count(), 1)
self.assertTrue(task.model_attempts.get().response_summary["outcome_unknown"])
self.assertEqual(self.direct_mock.synthesize.call_count, 1)
self.assertFalse(fallback.synthesize.called)
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)
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(
@@ -0,0 +1,407 @@
"""恢复被旧 ``generation_timeout`` 误判失败的异步视频任务。
恢复只消费已经存在的 Provider 任务 ID不重新提交模型旧任务的预留已释放
因此恢复结果由平台承担本次回归成本用户积分不追扣不重复预留
"""
from __future__ import annotations
from decimal import Decimal
from django.db import transaction
from django.utils import timezone
from apps.ai.models import AITask
from apps.assets.models import Asset
from apps.billing.models import CreditReservation
from apps.billing.pricing import quote_video_actual
from apps.projects.models import VideoSegment, VideoSegmentVersion
_RECOVERABLE_TIMEOUT_ERRORS = (
"视频生成超过配置的等待总时限",
"视频生成超过配置的成片等待总时限",
"生成超过配置的成片等待总时限",
)
_REMOTE_INFLIGHT = {"queued", "running", "processing", "submitted"}
_REMOTE_FAILED = {"failed", "expired", "cancelled"}
def _actual_video_model(task: AITask):
attempt = (
task.model_attempts.filter(status="succeeded", operation="video_generate")
.select_related("model_config__provider")
.order_by("-sequence")
.first()
)
return attempt.model_config if attempt and attempt.model_config else task.model_config
def _is_recoverable(task: AITask) -> bool:
payload = task.request_payload or {}
recovery = payload.get("timeout_recovery") or {}
if recovery.get("kind") == "generation_timeout_regression":
return True
return task.status == AITask.Status.FAILED and any(
marker in (task.error_message or "") for marker in _RECOVERABLE_TIMEOUT_ERRORS
)
def _safe_remote_summary(response: dict) -> dict:
usage = response.get("usage") or {}
return {
"status": str(response.get("status") or "").lower(),
"error": response.get("error") or None,
"total_tokens": usage.get("total_tokens"),
"has_content": bool(response.get("content")),
}
def _claim_recovery(task: AITask) -> tuple[AITask, str]:
with transaction.atomic():
locked = AITask.objects.select_for_update().get(id=task.id)
if locked.status == AITask.Status.SUCCEEDED:
return locked, "already_recovered"
if not _is_recoverable(locked):
raise ValueError("任务不是本次成片硬超时回归产生的可恢复失败任务")
if locked.credit_reservation.status != CreditReservation.Status.RELEASED:
raise ValueError("恢复任务的原预留并非已释放状态,拒绝绕过正常结算流程")
payload = dict(locked.request_payload or {})
recovery = dict(payload.get("timeout_recovery") or {})
recovery.update(
{
"kind": "generation_timeout_regression",
"state": "processing",
"original_error": recovery.get("original_error") or locked.error_message,
"user_charge": "waived",
}
)
payload["timeout_recovery"] = recovery
locked.request_payload = payload
locked.status = AITask.Status.POSTPROCESSING
locked.save(update_fields=["request_payload", "status", "updated_at"])
return locked, "claimed"
def _mark_recovery_failed(task_id, message: str) -> None:
with transaction.atomic():
locked = AITask.objects.select_for_update().get(id=task_id)
payload = dict(locked.request_payload or {})
recovery = dict(payload.get("timeout_recovery") or {})
recovery.update(
{
"kind": "generation_timeout_regression",
"state": "postprocess_failed",
"postprocess_error": message[:1000],
"user_charge": "waived",
}
)
payload["timeout_recovery"] = recovery
locked.request_payload = payload
locked.status = AITask.Status.FAILED
locked.save(update_fields=["request_payload", "status", "updated_at"])
def _platform_cost(response: dict, *, task: AITask, actual_model, with_video_ref: bool) -> tuple[int, Decimal]:
try:
total_tokens = int((response.get("usage") or {}).get("total_tokens") or 0)
except (TypeError, ValueError):
total_tokens = 0
if total_tokens <= 0:
return 0, task.base_cost
payload = task.request_payload or {}
settle = quote_video_actual(
actual_model,
tokens=total_tokens,
with_video_ref=with_video_ref,
resolution=str(payload.get("resolution") or "720p"),
multiplier=Decimal(str(payload.get("price_multiplier") or "1")),
)
return total_tokens, settle.base_cost_yuan
def _recover_project_video(*, task: AITask, actual_model, provider, response: dict) -> dict:
from apps.ai.services import _store_generated_media
segment_id = str((task.request_payload or {}).get("video_segment_id") or "")
segment = VideoSegment.objects.select_related("project").filter(id=segment_id, project=task.project).first()
if segment is None:
raise ValueError("任务对应的视频片段不存在")
# 同一片段可能因旧硬超时提示诱导用户重跑,导致多个远端任务后来都成功。
# 只允许其中一条恢复为平台版本,避免重复下载、重复资产和轮流覆盖采用结果。
for existing in segment.versions.select_related("task").order_by("created_at"):
existing_recovery = (
((existing.task.request_payload or {}).get("timeout_recovery") or {})
if existing.task
else {}
)
if (
existing.task_id != task.id
and existing_recovery.get("kind") == "generation_timeout_regression"
and existing_recovery.get("state") == "recovered"
):
return {
"action": "skipped_existing_recovery",
"task_id": str(task.id),
"video_segment_id": str(segment.id),
"existing_task_id": str(existing.task_id),
"existing_version_id": str(existing.id),
}
claimed, claim_state = _claim_recovery(task)
if claim_state == "already_recovered":
return {"action": "already_recovered", "task_id": str(task.id)}
try:
version = VideoSegmentVersion.objects.filter(task=claimed).select_related("asset").first()
if version is None:
asset = claimed.generated_assets.filter(
asset_type=Asset.Type.VIDEO,
purged_at__isnull=True,
).first()
if asset is None:
media = provider.extract_first_media_url(response)
asset = _store_generated_media(
team=segment.project.team,
user=claimed.created_by or segment.project.created_by,
project=segment.project,
task=claimed,
media=media,
name=f"{segment.project.name}-segment-{segment.sort_order + 1}",
category=Asset.Category.VIDEO_CLIP,
asset_type=Asset.Type.VIDEO,
)
total_tokens, base_cost = _platform_cost(
response,
task=claimed,
actual_model=actual_model,
with_video_ref=False,
)
with transaction.atomic():
locked = AITask.objects.select_for_update().get(id=claimed.id)
version = VideoSegmentVersion.objects.filter(task=locked).first()
if version is None:
version = VideoSegmentVersion.objects.create(
video_segment=segment,
task=locked,
asset=asset,
prompt=(locked.request_payload or {}).get("prompt", ""),
is_adopted=True,
metadata={"timeout_recovery": True, "user_charge": "waived"},
)
segment.versions.exclude(id=version.id).update(is_adopted=False)
if not version.is_adopted:
version.is_adopted = True
version.save(update_fields=["is_adopted", "updated_at"])
payload = dict(locked.request_payload or {})
recovery = dict(payload.get("timeout_recovery") or {})
recovery.update(
{
"state": "recovered",
"recovered_at": timezone.now().isoformat(),
"remote_total_tokens": total_tokens,
"user_charge": "waived",
}
)
payload["timeout_recovery"] = recovery
locked.status = AITask.Status.SUCCEEDED
locked.error_code = ""
locked.error_message = ""
locked.request_payload = payload
locked.response_payload = response
locked.actual_cost = Decimal("0")
locked.base_cost = base_cost
locked.completed_at = timezone.now()
locked.save(
update_fields=[
"status",
"error_code",
"error_message",
"request_payload",
"response_payload",
"actual_cost",
"base_cost",
"completed_at",
"updated_at",
]
)
segment.adopted_version = version
segment.status = VideoSegment.Status.SUCCEEDED
segment.error_message = ""
segment.save(update_fields=["adopted_version", "status", "error_message", "updated_at"])
return {
"action": "recovered",
"task_id": str(task.id),
"video_segment_id": str(segment.id),
"version_id": str(version.id),
"user_charge": "waived",
}
except Exception as exc:
_mark_recovery_failed(task.id, str(exc))
raise
def _recover_free_video(*, task: AITask, actual_model, provider, response: dict) -> dict:
from apps.ai.free_video import _store_free_video_media
claimed, claim_state = _claim_recovery(task)
if claim_state == "already_recovered":
return {"action": "already_recovered", "task_id": str(task.id)}
try:
asset = claimed.generated_assets.filter(
asset_type=Asset.Type.VIDEO,
purged_at__isnull=True,
).first()
if asset is None:
media = provider.extract_first_media_url(response)
asset = _store_free_video_media(task=claimed, media=media)
payload = dict(claimed.request_payload or {})
with_video_ref = any((item or {}).get("type") == "video" for item in payload.get("references") or [])
total_tokens, base_cost = _platform_cost(
response,
task=claimed,
actual_model=actual_model,
with_video_ref=with_video_ref,
)
with transaction.atomic():
locked = AITask.objects.select_for_update().get(id=claimed.id)
payload = dict(locked.request_payload or {})
recovery = dict(payload.get("timeout_recovery") or {})
recovery.update(
{
"state": "recovered",
"recovered_at": timezone.now().isoformat(),
"remote_total_tokens": total_tokens,
"asset_id": str(asset.id),
"user_charge": "waived",
}
)
payload["timeout_recovery"] = recovery
if total_tokens:
payload["actual_tokens"] = total_tokens
locked.status = AITask.Status.SUCCEEDED
locked.error_code = ""
locked.error_message = ""
locked.request_payload = payload
locked.response_payload = response
locked.actual_cost = Decimal("0")
locked.base_cost = base_cost
locked.completed_at = timezone.now()
locked.save(
update_fields=[
"status",
"error_code",
"error_message",
"request_payload",
"response_payload",
"actual_cost",
"base_cost",
"completed_at",
"updated_at",
]
)
return {
"action": "recovered",
"task_id": str(task.id),
"asset_id": str(asset.id),
"user_charge": "waived",
}
except Exception as exc:
_mark_recovery_failed(task.id, str(exc))
raise
def reconcile_video_timeout_task(task_id: str, *, apply: bool = False) -> dict:
"""查询远端状态;默认只读,``apply=True`` 时仅恢复已成功结果。"""
from apps.ai.services import get_video_provider
task = (
AITask.objects.select_related(
"model_config__provider",
"project",
"created_by",
"credit_reservation",
)
.filter(id=task_id)
.first()
)
if task is None:
raise ValueError("AI 任务不存在")
if task.task_type not in {AITask.Type.VIDEO_SEGMENT, AITask.Type.FREE_VIDEO}:
raise ValueError("只支持视频片段或自由创作视频任务")
if task.status == AITask.Status.SUCCEEDED and (task.request_payload or {}).get("timeout_recovery"):
return {
"task_id": str(task.id),
"task_type": task.task_type,
"action": "already_recovered",
"applied": False,
}
if not _is_recoverable(task):
raise ValueError("任务不是本次成片硬超时回归产生的可恢复失败任务")
if not task.provider_task_id:
raise ValueError("任务没有 Provider 任务 ID,禁止重新提交模型")
actual_model = _actual_video_model(task)
provider = get_video_provider(actual_model)
response = provider.poll_video_task(
endpoint=actual_model.endpoint,
provider_task_id=task.provider_task_id,
timeout=load_poll_request_timeout(),
)
summary = _safe_remote_summary(response)
result = {
"task_id": str(task.id),
"task_type": task.task_type,
"model": actual_model.name,
"provider": actual_model.provider.name,
"remote": summary,
"action": "preview",
"applied": False,
}
if not apply:
return result
remote_status = summary["status"]
if remote_status in _REMOTE_INFLIGHT:
result["action"] = "still_running"
return result
if remote_status in _REMOTE_FAILED:
result["action"] = "remote_failed"
return result
if remote_status != "succeeded":
result["action"] = "unknown_remote_status"
return result
if task.task_type == AITask.Type.VIDEO_SEGMENT:
applied = _recover_project_video(
task=task,
actual_model=actual_model,
provider=provider,
response=response,
)
else:
applied = _recover_free_video(
task=task,
actual_model=actual_model,
provider=provider,
response=response,
)
result.update(applied)
result["applied"] = result.get("action") == "recovered"
return result
def load_poll_request_timeout() -> float:
from apps.ai.routing_policy import load_model_routing_policy
return load_model_routing_policy().video.poll_request_timeout
__all__ = ["reconcile_video_timeout_task"]