fix: 修正模型超时与结果未知重放
This commit is contained in:
@@ -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")
|
||||
|
||||
@@ -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))
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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"))
|
||||
|
||||
@@ -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"))
|
||||
|
||||
@@ -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"]
|
||||
Reference in New Issue
Block a user