feat: 接入模型动态 Fallback 与调用审计
This commit is contained in:
@@ -1,7 +1,18 @@
|
||||
from django.apps import AppConfig
|
||||
from django.core.exceptions import ImproperlyConfigured
|
||||
|
||||
|
||||
class AiConfig(AppConfig):
|
||||
default_auto_field = "django.db.models.BigAutoField"
|
||||
name = "apps.ai"
|
||||
|
||||
def ready(self) -> None:
|
||||
# API 与 Celery Worker 都会在 Django 启动时经过这里。路由策略一旦配置错误就立即失败,
|
||||
# 避免两端带着不同或不可控的超时、重试规则继续运行。
|
||||
from .routing_policy import RoutingPolicyConfigurationError, load_model_routing_policy
|
||||
|
||||
try:
|
||||
load_model_routing_policy()
|
||||
except RoutingPolicyConfigurationError as exc:
|
||||
raise ImproperlyConfigured(f"模型路由策略配置错误:{exc}") from exc
|
||||
|
||||
|
||||
@@ -2,12 +2,15 @@ VOLCANO_PROVIDER = {
|
||||
"name": "volcengine",
|
||||
"display_name": "火山引擎(豆包)",
|
||||
"base_url": "https://ark.cn-beijing.volces.com/api/v3",
|
||||
# Fallback 候选排序:数字越小越优先。运行时只读配置,不识别供应商名称。
|
||||
"metadata": {"routing": {"fallback_priority": 10}},
|
||||
}
|
||||
|
||||
YUNQI_PROVIDER = {
|
||||
"name": "yunqi",
|
||||
"display_name": "YunQi AI(New API 网关)",
|
||||
"base_url": "https://www.yunqiai.chat/v1",
|
||||
"metadata": {"routing": {"fallback_priority": 20}},
|
||||
}
|
||||
|
||||
YUNQI_MODELS = [
|
||||
@@ -16,7 +19,12 @@ YUNQI_MODELS = [
|
||||
"name": "gpt-image-2",
|
||||
"capability": "image",
|
||||
"endpoint": "images/generations",
|
||||
"metadata": {"modes": ["text"], "response_format": "b64_json", "source": "https://www.yunqiai.chat"},
|
||||
"metadata": {
|
||||
"modes": ["text", "singleImage", "multiReference"],
|
||||
"supports_reference": True,
|
||||
"response_format": "b64_json",
|
||||
"source": "https://www.yunqiai.chat",
|
||||
},
|
||||
},
|
||||
]
|
||||
|
||||
@@ -132,3 +140,78 @@ VOLCANO_MODELS = [
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
IMAGE_RATIOS = ["1:1", "3:4", "4:5", "9:16", "16:9"]
|
||||
VIDEO_RATIOS = ["21:9", "16:9", "4:3", "1:1", "3:4", "9:16"]
|
||||
|
||||
|
||||
def _mode_limit(modes: set[str], prefix: str) -> int:
|
||||
for mode in modes:
|
||||
if mode.startswith(f"{prefix}:"):
|
||||
try:
|
||||
return int(mode.split(":", 1)[1])
|
||||
except (TypeError, ValueError):
|
||||
return 0
|
||||
return 0
|
||||
|
||||
|
||||
def _catalog_capabilities(item: dict) -> dict:
|
||||
"""由模型目录声明生成最小能力契约;只用于写库初始化,不参与运行时猜测。"""
|
||||
|
||||
capability = item["capability"]
|
||||
metadata = item["metadata"]
|
||||
if capability == "text":
|
||||
return {"operations": ["chat"], "features": ["streaming", "structured_output"]}
|
||||
if capability == "image":
|
||||
modes = set(metadata.get("modes") or [])
|
||||
supports_reference = bool(metadata.get("supports_reference")) or bool(
|
||||
modes & {"singleImage", "multiReference"}
|
||||
)
|
||||
return {
|
||||
"operations": ["image_generate", "image_edit"] if supports_reference else ["image_generate"],
|
||||
"features": [],
|
||||
"reference_modes": ["none", "single", "multiple"] if supports_reference else ["none"],
|
||||
"max_reference_images": 9 if supports_reference else 0,
|
||||
"aspect_ratios": IMAGE_RATIOS,
|
||||
}
|
||||
if capability == "video":
|
||||
modes = set(metadata.get("modes") or [])
|
||||
features = ["text_to_video"]
|
||||
for mode, feature in (
|
||||
("startFrameOptional", "start_frame"),
|
||||
("lastFrameOptional", "last_frame"),
|
||||
("imageReference:9", "image_reference"),
|
||||
("videoReference:3", "video_reference"),
|
||||
("audioReference:3", "audio_reference"),
|
||||
):
|
||||
if mode in modes:
|
||||
features.append(feature)
|
||||
if metadata.get("audio") in {"optional", "required", True}:
|
||||
features.append("generate_audio")
|
||||
return {
|
||||
"operations": ["video_generate"],
|
||||
"features": features,
|
||||
"max_reference_images": _mode_limit(modes, "imageReference"),
|
||||
"max_reference_videos": _mode_limit(modes, "videoReference"),
|
||||
"max_reference_audios": _mode_limit(modes, "audioReference"),
|
||||
"aspect_ratios": VIDEO_RATIOS,
|
||||
"resolutions": list(metadata.get("resolutions") or []),
|
||||
"durations": list(metadata.get("durations") or []),
|
||||
}
|
||||
return {}
|
||||
|
||||
|
||||
def _enrich_catalog_routing(items: list[dict], *, fallback_on_failure: bool) -> None:
|
||||
for item in items:
|
||||
metadata = item["metadata"]
|
||||
metadata["routing"] = {
|
||||
"fallback_on_failure": fallback_on_failure,
|
||||
"fallback_candidate": True,
|
||||
}
|
||||
metadata["capabilities"] = _catalog_capabilities(item)
|
||||
|
||||
|
||||
# 火山当前只作为候选、不向外切;YunQi 可在失败后切换。以后调整只改数据库 metadata。
|
||||
_enrich_catalog_routing(VOLCANO_MODELS, fallback_on_failure=False)
|
||||
_enrich_catalog_routing(YUNQI_MODELS, fallback_on_failure=True)
|
||||
|
||||
|
||||
@@ -15,6 +15,7 @@ import logging
|
||||
import re
|
||||
import uuid
|
||||
from datetime import timedelta
|
||||
from decimal import Decimal
|
||||
from io import BytesIO
|
||||
|
||||
from django.conf import settings
|
||||
@@ -272,18 +273,33 @@ def _reap_stale_free_video_tasks(*, team) -> None:
|
||||
· SUBMITTED/POLLING 超 2 小时:轮询链早已断且无人认领(正常出片 5-10 分钟)→ 标失败退费;
|
||||
· POSTPROCESSING 超 30 分钟:转存/结算中途崩溃 → 标失败退费(火山可能已出片,平台承担该笔成本)。"""
|
||||
now = timezone.now()
|
||||
from .routing_policy import load_model_routing_policy
|
||||
|
||||
video_policy = load_model_routing_policy().video
|
||||
buckets = [
|
||||
([AITask.Status.RESERVED], now - timedelta(minutes=10), "任务未在预期时间内提交(自动回收)"),
|
||||
([AITask.Status.SUBMITTED, AITask.Status.POLLING], now - timedelta(hours=2), "生成超时(自动回收)"),
|
||||
([AITask.Status.POSTPROCESSING], now - timedelta(minutes=30), "视频结果处理超时(自动回收)"),
|
||||
(
|
||||
[AITask.Status.RESERVED],
|
||||
{"updated_at__lt": now - timedelta(seconds=video_policy.submit_total_timeout)},
|
||||
"任务未在配置的提交总时限内完成(自动回收)",
|
||||
),
|
||||
(
|
||||
[AITask.Status.SUBMITTED, AITask.Status.POLLING],
|
||||
{"submitted_at__lt": now - timedelta(seconds=video_policy.generation_timeout)},
|
||||
"生成超过配置的成片等待总时限(自动回收)",
|
||||
),
|
||||
(
|
||||
[AITask.Status.POSTPROCESSING],
|
||||
{"updated_at__lt": now - timedelta(minutes=30)},
|
||||
"视频结果处理超时(自动回收)",
|
||||
),
|
||||
]
|
||||
for statuses, cutoff, reason in buckets:
|
||||
for statuses, stale_filter, reason in buckets:
|
||||
stale = AITask.objects.filter(
|
||||
team=team,
|
||||
project__isnull=True,
|
||||
task_type=AITask.Type.FREE_VIDEO,
|
||||
status__in=statuses,
|
||||
updated_at__lt=cutoff,
|
||||
**stale_filter,
|
||||
)
|
||||
for task in stale:
|
||||
try:
|
||||
@@ -404,6 +420,7 @@ def submit_free_video(*, team, user, params: dict) -> AITask:
|
||||
"price_multiplier": quote.meta.get("price_multiplier", "1"),
|
||||
"points_per_yuan_snapshot": quote.meta.get("rate", ""),
|
||||
"references": built["snapshots"],
|
||||
"model_routing_v1": True,
|
||||
}
|
||||
|
||||
# 建任务 + 预留同一事务:余额不足/限额拦截时回滚任务行,不留半套
|
||||
@@ -418,7 +435,7 @@ def submit_free_video(*, team, user, params: dict) -> AITask:
|
||||
idempotency_key=f"free_video:{team.id}:{uuid.uuid4()}",
|
||||
request_payload=request_payload,
|
||||
estimated_cost=quote.points,
|
||||
base_cost=quote.base_cost_yuan,
|
||||
base_cost=Decimal("0"),
|
||||
)
|
||||
try:
|
||||
reserve_credit(team=team, user=user, task=task, amount=reserve_amount)
|
||||
@@ -431,26 +448,44 @@ def submit_free_video(*, team, user, params: dict) -> AITask:
|
||||
|
||||
# 火山调用在事务外(不持锁调外网)
|
||||
try:
|
||||
from .services import build_provider
|
||||
from .services import execute_routed_video_submit
|
||||
|
||||
provider = build_provider(model_config)
|
||||
response = provider.create_video_task(
|
||||
model=model_config.name,
|
||||
endpoint=model_config.endpoint,
|
||||
routed = execute_routed_video_submit(
|
||||
task=task,
|
||||
primary_model=model_config,
|
||||
prompt=built["api_prompt"],
|
||||
ratio=aspect_ratio,
|
||||
duration=duration,
|
||||
resolution=resolution,
|
||||
generate_audio=generate_audio,
|
||||
reference_images=[],
|
||||
content_items=built["content_items"],
|
||||
pricing_references=built["snapshots"],
|
||||
generate_audio=generate_audio,
|
||||
seed=seed if seed != -1 else None,
|
||||
search_mode=search_mode,
|
||||
request_summary={"feature": "free_video", "mode": mode},
|
||||
)
|
||||
task.provider_task_id = str(response.get("id") or response.get("task_id") or "")
|
||||
response, provider_task_id = routed.value
|
||||
# execute_model_call 以原子 F 表达式累计实际尝试的平台成本;刷新内存对象,
|
||||
# 保证本方法返回值与数据库中的 AITask.base_cost 完全一致。
|
||||
task.refresh_from_db(fields=["base_cost"])
|
||||
task.provider_task_id = provider_task_id
|
||||
task.response_payload = response
|
||||
payload = dict(task.request_payload or {})
|
||||
payload["actual_model_config_id"] = str(routed.actual_model.id)
|
||||
task.request_payload = payload
|
||||
task.status = AITask.Status.SUBMITTED
|
||||
task.submitted_at = timezone.now()
|
||||
task.save(update_fields=["provider_task_id", "response_payload", "status", "submitted_at", "updated_at"])
|
||||
task.save(
|
||||
update_fields=[
|
||||
"provider_task_id",
|
||||
"response_payload",
|
||||
"request_payload",
|
||||
"status",
|
||||
"submitted_at",
|
||||
"updated_at",
|
||||
]
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001 — 创建失败:标失败退费,返回失败卡(不向上抛)
|
||||
code, raw_message = parse_provider_error(exc)
|
||||
public_error = classify_generation_error(
|
||||
@@ -577,11 +612,47 @@ def finalize_free_video(*, task: AITask) -> AITask:
|
||||
if not task.provider_task_id:
|
||||
return task
|
||||
|
||||
from .services import build_provider
|
||||
from .routing_policy import load_model_routing_policy
|
||||
from .services import get_video_provider
|
||||
|
||||
provider = build_provider(task.model_config)
|
||||
video_policy = load_model_routing_policy().video
|
||||
if task.submitted_at and (
|
||||
timezone.now() - task.submitted_at
|
||||
).total_seconds() >= video_policy.generation_timeout:
|
||||
timeout_message = "视频生成超过配置的成片等待总时限"
|
||||
public_error = classify_generation_error(
|
||||
TimeoutError(timeout_message),
|
||||
operation="video_generate",
|
||||
reference_id=str(task.id),
|
||||
)
|
||||
with transaction.atomic():
|
||||
locked = AITask.objects.select_for_update().get(id=task.id)
|
||||
if locked.status not in (AITask.Status.SUBMITTED, AITask.Status.POLLING):
|
||||
return locked
|
||||
locked.status = AITask.Status.FAILED
|
||||
locked.error_code = "GenerationTimeout"
|
||||
locked.error_message = timeout_message
|
||||
locked.completed_at = timezone.now()
|
||||
locked.save(
|
||||
update_fields=["status", "error_code", "error_message", "completed_at", "updated_at"]
|
||||
)
|
||||
release_credit(reservation=locked.credit_reservation, reason=timeout_message)
|
||||
_notify_failure(locked, raw=timeout_message, hint=public_error.fallback_message)
|
||||
return locked
|
||||
|
||||
submit_attempt = (
|
||||
task.model_attempts.filter(status="succeeded", operation="video_generate")
|
||||
.select_related("model_config__provider")
|
||||
.order_by("-sequence")
|
||||
.first()
|
||||
)
|
||||
actual_model = submit_attempt.model_config if submit_attempt and submit_attempt.model_config else task.model_config
|
||||
|
||||
provider = get_video_provider(actual_model)
|
||||
response = provider.poll_video_task(
|
||||
endpoint=task.model_config.endpoint, provider_task_id=task.provider_task_id
|
||||
endpoint=actual_model.endpoint,
|
||||
provider_task_id=task.provider_task_id,
|
||||
timeout=video_policy.poll_request_timeout,
|
||||
)
|
||||
remote_status = str(response.get("status") or "")
|
||||
|
||||
@@ -647,7 +718,7 @@ def finalize_free_video(*, task: AITask) -> AITask:
|
||||
from decimal import Decimal
|
||||
|
||||
settle = quote_video_actual(
|
||||
locked.model_config, tokens=total_tokens, with_video_ref=with_video_ref, resolution=resolution,
|
||||
actual_model, tokens=total_tokens, with_video_ref=with_video_ref, resolution=resolution,
|
||||
multiplier=Decimal(str(payload.get("price_multiplier") or "1")),
|
||||
)
|
||||
actual, base_cost = settle.points, settle.base_cost_yuan
|
||||
|
||||
@@ -0,0 +1,154 @@
|
||||
"""为存量供应商和模型补齐动态 Fallback 路由与最小能力配置。
|
||||
|
||||
运行时解析器不识别供应商名或模型名;这里仅把当前产品已经确认的火山/YunQi优先级
|
||||
写入数据库 metadata。以后新增供应商或模型只维护 metadata,不需要修改路由代码。
|
||||
"""
|
||||
|
||||
from django.db import migrations
|
||||
|
||||
|
||||
DIRECT_PROVIDER_NAMES = {"volcengine", "volcano", "ark", "volcano_ark", "doubao"}
|
||||
IMAGE_RATIOS = ["1:1", "3:4", "4:5", "9:16", "16:9"]
|
||||
VIDEO_RATIOS = ["21:9", "16:9", "4:3", "1:1", "3:4", "9:16"]
|
||||
VOICE_MAP = {
|
||||
"BV700_streaming": "BV700_streaming",
|
||||
"BV034_streaming": "BV034_streaming",
|
||||
"BV001_streaming": "BV001_streaming",
|
||||
"BV056_streaming": "BV056_streaming",
|
||||
"BV102_streaming": "BV102_streaming",
|
||||
"BV002_streaming": "BV002_streaming",
|
||||
}
|
||||
|
||||
|
||||
def _provider_priority(name):
|
||||
if name in DIRECT_PROVIDER_NAMES:
|
||||
return 10
|
||||
if name == "yunqi" or name.startswith("yunqi_"):
|
||||
return 20
|
||||
return 100
|
||||
|
||||
|
||||
def _mode_limit(modes, prefix, default=0):
|
||||
for mode in modes:
|
||||
if mode.startswith(prefix):
|
||||
_, _, raw = mode.partition(":")
|
||||
try:
|
||||
return int(raw)
|
||||
except (TypeError, ValueError):
|
||||
return default
|
||||
return default
|
||||
|
||||
|
||||
def _capabilities(model):
|
||||
metadata = dict(model.metadata or {})
|
||||
capability = model.capability
|
||||
if capability == "text":
|
||||
return {
|
||||
"operations": ["chat"],
|
||||
"features": ["streaming", "structured_output"],
|
||||
}
|
||||
if capability == "image":
|
||||
modes = {str(item) for item in metadata.get("modes") or []}
|
||||
supports_reference = bool(metadata.get("supports_reference")) or bool(
|
||||
modes & {"singleImage", "multiReference"}
|
||||
)
|
||||
reference_modes = ["none"]
|
||||
if supports_reference or "singleImage" in modes:
|
||||
reference_modes.append("single")
|
||||
if supports_reference or "multiReference" in modes:
|
||||
reference_modes.append("multiple")
|
||||
operations = ["image_generate"]
|
||||
if supports_reference:
|
||||
operations.append("image_edit")
|
||||
return {
|
||||
"operations": operations,
|
||||
"features": [],
|
||||
"reference_modes": reference_modes,
|
||||
"max_reference_images": 9 if supports_reference else 0,
|
||||
"aspect_ratios": IMAGE_RATIOS,
|
||||
}
|
||||
if capability == "video":
|
||||
modes = {str(item) for item in metadata.get("modes") or []}
|
||||
features = ["text_to_video"]
|
||||
for mode, feature in (
|
||||
("startFrameOptional", "start_frame"),
|
||||
("lastFrameOptional", "last_frame"),
|
||||
("imageReference:9", "image_reference"),
|
||||
("videoReference:3", "video_reference"),
|
||||
("audioReference:3", "audio_reference"),
|
||||
):
|
||||
if mode in modes:
|
||||
features.append(feature)
|
||||
if metadata.get("audio") in {"optional", "required", True}:
|
||||
features.append("generate_audio")
|
||||
return {
|
||||
"operations": ["video_generate"],
|
||||
"features": features,
|
||||
"max_reference_images": _mode_limit(modes, "imageReference", 0),
|
||||
"max_reference_videos": _mode_limit(modes, "videoReference", 0),
|
||||
"max_reference_audios": _mode_limit(modes, "audioReference", 0),
|
||||
"aspect_ratios": VIDEO_RATIOS,
|
||||
"resolutions": list(metadata.get("resolutions") or []),
|
||||
"durations": list(metadata.get("durations") or []),
|
||||
}
|
||||
if capability == "audio":
|
||||
return {
|
||||
"operations": ["tts"],
|
||||
"features": [],
|
||||
"languages": ["zh-CN"],
|
||||
"voice_map": VOICE_MAP,
|
||||
"max_chars": 10000,
|
||||
"speed_range": [0.5, 2.0],
|
||||
"output_formats": ["mp3"],
|
||||
}
|
||||
return {}
|
||||
|
||||
|
||||
def apply(apps, schema_editor):
|
||||
ModelProvider = apps.get_model("ai", "ModelProvider")
|
||||
ModelConfig = apps.get_model("ai", "ModelConfig")
|
||||
|
||||
for provider in ModelProvider.objects.all().iterator():
|
||||
metadata = dict(provider.metadata or {})
|
||||
routing = dict(metadata.get("routing") or {})
|
||||
routing.setdefault("fallback_priority", _provider_priority(provider.name))
|
||||
metadata["routing"] = routing
|
||||
provider.metadata = metadata
|
||||
provider.save(update_fields=["metadata"])
|
||||
|
||||
for model in ModelConfig.objects.select_related("provider").all().iterator():
|
||||
capabilities = _capabilities(model)
|
||||
metadata = dict(model.metadata or {})
|
||||
routing = dict(metadata.get("routing") or {})
|
||||
is_supported = bool(capabilities)
|
||||
routing.setdefault("fallback_candidate", is_supported)
|
||||
routing.setdefault(
|
||||
"fallback_on_failure",
|
||||
is_supported and model.provider.name not in DIRECT_PROVIDER_NAMES,
|
||||
)
|
||||
metadata["routing"] = routing
|
||||
if capabilities:
|
||||
metadata.setdefault("capabilities", capabilities)
|
||||
model.metadata = metadata
|
||||
model.save(update_fields=["metadata"])
|
||||
|
||||
|
||||
def revert(apps, schema_editor):
|
||||
ModelProvider = apps.get_model("ai", "ModelProvider")
|
||||
ModelConfig = apps.get_model("ai", "ModelConfig")
|
||||
for provider in ModelProvider.objects.all().iterator():
|
||||
metadata = dict(provider.metadata or {})
|
||||
metadata.pop("routing", None)
|
||||
provider.metadata = metadata
|
||||
provider.save(update_fields=["metadata"])
|
||||
for model in ModelConfig.objects.all().iterator():
|
||||
metadata = dict(model.metadata or {})
|
||||
metadata.pop("routing", None)
|
||||
metadata.pop("capabilities", None)
|
||||
model.metadata = metadata
|
||||
model.save(update_fields=["metadata"])
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
dependencies = [("ai", "0027_aitask_model_triview_type")]
|
||||
operations = [migrations.RunPython(apply, revert)]
|
||||
@@ -0,0 +1,91 @@
|
||||
from django.db import migrations, models
|
||||
import django.db.models.deletion
|
||||
import uuid
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
dependencies = [("ai", "0028_seed_model_routing_metadata")]
|
||||
|
||||
operations = [
|
||||
migrations.CreateModel(
|
||||
name="AIModelAttempt",
|
||||
fields=[
|
||||
("id", models.UUIDField(default=uuid.uuid4, editable=False, primary_key=True, serialize=False)),
|
||||
("created_at", models.DateTimeField(auto_now_add=True)),
|
||||
("updated_at", models.DateTimeField(auto_now=True)),
|
||||
("sequence", models.PositiveIntegerField()),
|
||||
("provider_name", models.CharField(max_length=128)),
|
||||
("provider_display_name", models.CharField(blank=True, max_length=128)),
|
||||
("model_name", models.CharField(max_length=128)),
|
||||
("model_display_name", models.CharField(blank=True, max_length=128)),
|
||||
("public_model_name", models.CharField(blank=True, max_length=128)),
|
||||
("capability", models.CharField(max_length=32)),
|
||||
("operation", models.CharField(max_length=64)),
|
||||
(
|
||||
"status",
|
||||
models.CharField(
|
||||
choices=[("started", "Started"), ("succeeded", "Succeeded"), ("failed", "Failed")],
|
||||
default="started",
|
||||
max_length=24,
|
||||
),
|
||||
),
|
||||
("is_retry", models.BooleanField(default=False)),
|
||||
("is_fallback", models.BooleanField(default=False)),
|
||||
("provider_task_id", models.CharField(blank=True, max_length=255)),
|
||||
("started_at", models.DateTimeField()),
|
||||
("finished_at", models.DateTimeField(blank=True, null=True)),
|
||||
("duration_ms", models.PositiveBigIntegerField(blank=True, null=True)),
|
||||
("error_type", models.CharField(blank=True, max_length=64)),
|
||||
("provider_error_code", models.CharField(blank=True, max_length=128)),
|
||||
("raw_error", models.TextField(blank=True)),
|
||||
("safe_error_summary", models.TextField(blank=True)),
|
||||
("usage", models.JSONField(blank=True, default=dict)),
|
||||
("platform_cost", models.DecimalField(decimal_places=4, default=0, max_digits=12)),
|
||||
("request_summary", models.JSONField(blank=True, default=dict)),
|
||||
("response_summary", models.JSONField(blank=True, default=dict)),
|
||||
(
|
||||
"model_config",
|
||||
models.ForeignKey(
|
||||
blank=True,
|
||||
null=True,
|
||||
on_delete=django.db.models.deletion.SET_NULL,
|
||||
to="ai.modelconfig",
|
||||
),
|
||||
),
|
||||
(
|
||||
"previous_attempt",
|
||||
models.ForeignKey(
|
||||
blank=True,
|
||||
null=True,
|
||||
on_delete=django.db.models.deletion.SET_NULL,
|
||||
related_name="next_attempts",
|
||||
to="ai.aimodelattempt",
|
||||
),
|
||||
),
|
||||
(
|
||||
"provider",
|
||||
models.ForeignKey(
|
||||
blank=True,
|
||||
null=True,
|
||||
on_delete=django.db.models.deletion.SET_NULL,
|
||||
to="ai.modelprovider",
|
||||
),
|
||||
),
|
||||
(
|
||||
"task",
|
||||
models.ForeignKey(
|
||||
on_delete=django.db.models.deletion.CASCADE,
|
||||
related_name="model_attempts",
|
||||
to="ai.aitask",
|
||||
),
|
||||
),
|
||||
],
|
||||
options={
|
||||
"ordering": ["sequence"],
|
||||
"indexes": [models.Index(fields=["task", "status"], name="ai_attempt_task_status_idx")],
|
||||
"constraints": [
|
||||
models.UniqueConstraint(fields=("task", "sequence"), name="ai_attempt_task_sequence_unique")
|
||||
],
|
||||
},
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,329 @@
|
||||
"""模型能力契约与动态 Fallback 候选解析。
|
||||
|
||||
本模块只负责“某模型是否满足本次调用”与“失败后候选如何排序”,不执行 Provider
|
||||
请求、重试、日志或账务。业务入口后续只声明 :class:`ModelRequirements`,不得按模型名
|
||||
复制能力判断。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from apps.ai.models import ModelConfig, ModelProvider
|
||||
from apps.ai.routing_policy import load_model_routing_policy
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ModelRequirements:
|
||||
"""一次模型调用的最小能力需求;未使用的字段保持空值。"""
|
||||
|
||||
capability: str
|
||||
operation: str
|
||||
features: frozenset[str] = field(default_factory=frozenset)
|
||||
reference_mode: str | None = None
|
||||
reference_images: int = 0
|
||||
reference_videos: int = 0
|
||||
reference_audios: int = 0
|
||||
aspect_ratio: str | None = None
|
||||
resolution: str | None = None
|
||||
duration: int | None = None
|
||||
language: str | None = None
|
||||
public_voice: str | None = None
|
||||
char_count: int | None = None
|
||||
speed_ratio: float | None = None
|
||||
output_format: str | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if not self.capability or not self.operation:
|
||||
raise ValueError("模型能力需求必须包含 capability 和 operation")
|
||||
for name in ("reference_images", "reference_videos", "reference_audios"):
|
||||
if getattr(self, name) < 0:
|
||||
raise ValueError(f"{name} 不能为负数")
|
||||
if self.reference_mode not in {None, "none", "single", "multiple"}:
|
||||
raise ValueError("reference_mode 只能是 none、single 或 multiple")
|
||||
if self.reference_mode == "single" and self.reference_images != 1:
|
||||
raise ValueError("reference_mode=single 时 reference_images 必须等于 1")
|
||||
if self.reference_mode == "multiple" and self.reference_images < 2:
|
||||
raise ValueError("reference_mode=multiple 时 reference_images 必须至少为 2")
|
||||
if self.duration is not None and self.duration <= 0:
|
||||
raise ValueError("duration 必须大于 0")
|
||||
if self.char_count is not None and self.char_count < 0:
|
||||
raise ValueError("char_count 不能为负数")
|
||||
if self.speed_ratio is not None and self.speed_ratio <= 0:
|
||||
raise ValueError("speed_ratio 必须大于 0")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CapabilityMatch:
|
||||
matched: bool
|
||||
reasons: tuple[str, ...] = ()
|
||||
|
||||
|
||||
def _dict(value: Any) -> dict[str, Any]:
|
||||
return value if isinstance(value, dict) else {}
|
||||
|
||||
|
||||
def _set(value: Any) -> set[str]:
|
||||
if not isinstance(value, (list, tuple, set, frozenset)):
|
||||
return set()
|
||||
return {str(item) for item in value if str(item)}
|
||||
|
||||
|
||||
def _positive_int(value: Any) -> int | None:
|
||||
if isinstance(value, bool):
|
||||
return None
|
||||
try:
|
||||
result = int(value)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
return result if result >= 0 else None
|
||||
|
||||
|
||||
def routing_metadata(model: ModelConfig) -> dict[str, Any]:
|
||||
return _dict(_dict(model.metadata).get("routing"))
|
||||
|
||||
|
||||
def capability_metadata(model: ModelConfig) -> dict[str, Any]:
|
||||
return _dict(_dict(model.metadata).get("capabilities"))
|
||||
|
||||
|
||||
def model_allows_fallback(model: ModelConfig) -> bool:
|
||||
"""该模型失败后是否允许向外切换;缺失配置时保守关闭。"""
|
||||
|
||||
return routing_metadata(model).get("fallback_on_failure") is True
|
||||
|
||||
|
||||
def model_is_fallback_candidate(model: ModelConfig) -> bool:
|
||||
"""该模型是否允许被其他失败任务选中;缺失配置时保守关闭。"""
|
||||
|
||||
return routing_metadata(model).get("fallback_candidate") is True
|
||||
|
||||
|
||||
def provider_fallback_priority(provider: ModelProvider) -> int:
|
||||
"""读取供应商候选优先级;数字越小越优先,缺失或非法值统一排到普通供应商层。"""
|
||||
|
||||
value = _dict(_dict(provider.metadata).get("routing")).get("fallback_priority", 100)
|
||||
if isinstance(value, bool):
|
||||
return 100
|
||||
try:
|
||||
return int(value)
|
||||
except (TypeError, ValueError):
|
||||
return 100
|
||||
|
||||
|
||||
def provider_metadata_errors(metadata: Any) -> tuple[str, ...]:
|
||||
"""校验供应商路由配置;供后台写入校验和运行前诊断共同复用。"""
|
||||
|
||||
routing = _dict(_dict(metadata).get("routing"))
|
||||
if "fallback_priority" not in routing:
|
||||
return ()
|
||||
value = routing["fallback_priority"]
|
||||
if isinstance(value, bool) or not isinstance(value, int) or not 0 <= value <= 1000:
|
||||
return ("routing.fallback_priority 必须是 0 到 1000 的整数,数字越小越优先",)
|
||||
return ()
|
||||
|
||||
|
||||
def model_metadata_errors(capability: str, metadata: Any) -> tuple[str, ...]:
|
||||
"""校验模型路由开关与最小能力契约,缺少候选开关时允许保存。"""
|
||||
|
||||
root = _dict(metadata)
|
||||
routing = _dict(root.get("routing"))
|
||||
errors: list[str] = []
|
||||
for name in ("fallback_on_failure", "fallback_candidate"):
|
||||
if name in routing and not isinstance(routing[name], bool):
|
||||
errors.append(f"routing.{name} 必须是 true 或 false")
|
||||
|
||||
if routing.get("fallback_candidate") is not True:
|
||||
return tuple(errors)
|
||||
|
||||
capabilities = _dict(root.get("capabilities"))
|
||||
if not _set(capabilities.get("operations")):
|
||||
errors.append("capabilities.operations 至少配置一个操作")
|
||||
if capability == ModelConfig.Capability.IMAGE:
|
||||
modes = _set(capabilities.get("reference_modes"))
|
||||
if modes - {"none", "single", "multiple"}:
|
||||
errors.append("capabilities.reference_modes 只能包含 none、single、multiple")
|
||||
if capability == ModelConfig.Capability.VIDEO:
|
||||
if not _set(capabilities.get("resolutions")):
|
||||
errors.append("视频候选必须配置 capabilities.resolutions")
|
||||
durations = capabilities.get("durations")
|
||||
if not isinstance(durations, (list, tuple)) or not durations:
|
||||
errors.append("视频候选必须配置 capabilities.durations")
|
||||
if capability == ModelConfig.Capability.AUDIO:
|
||||
voice_map = capabilities.get("voice_map")
|
||||
if voice_map is not None and not isinstance(voice_map, dict):
|
||||
errors.append("配音 capabilities.voice_map 必须是公开音色到供应商音色 ID 的对象")
|
||||
return tuple(errors)
|
||||
|
||||
|
||||
def capability_metadata_errors(model: ModelConfig) -> tuple[str, ...]:
|
||||
"""返回会导致模型无法安全进入候选池的中文配置问题。"""
|
||||
|
||||
return model_metadata_errors(model.capability, model.metadata)
|
||||
|
||||
|
||||
def _check_limit(
|
||||
reasons: list[str],
|
||||
capabilities: dict[str, Any],
|
||||
field_name: str,
|
||||
required: int,
|
||||
label: str,
|
||||
) -> None:
|
||||
if required <= 0:
|
||||
return
|
||||
maximum = _positive_int(capabilities.get(field_name))
|
||||
if maximum is None or maximum < required:
|
||||
reasons.append(f"{label}上限不足:需要 {required},配置为 {maximum if maximum is not None else '缺失'}")
|
||||
|
||||
|
||||
def _video_pricing_available(model: ModelConfig, resolution: str | None) -> bool:
|
||||
if not resolution:
|
||||
return True
|
||||
pricing = _dict(_dict(model.metadata).get("pricing"))
|
||||
if not pricing:
|
||||
return False
|
||||
# 1080p/4k 必须有精确价格档;480p/720p 可使用现有 default 档。
|
||||
tier = pricing.get(resolution) if resolution in {"1080p", "4k"} else pricing.get(resolution) or pricing.get("default")
|
||||
return isinstance(tier, dict) and bool(tier)
|
||||
|
||||
|
||||
def match_model_requirements(model: ModelConfig, requirements: ModelRequirements) -> CapabilityMatch:
|
||||
"""纯函数式能力匹配;缺失配置一律保守排除,不按模型名猜测。"""
|
||||
|
||||
reasons: list[str] = []
|
||||
if model.capability != requirements.capability:
|
||||
reasons.append(f"能力大类不匹配:需要 {requirements.capability},模型为 {model.capability}")
|
||||
return CapabilityMatch(False, tuple(reasons))
|
||||
|
||||
capabilities = capability_metadata(model)
|
||||
operations = _set(capabilities.get("operations"))
|
||||
if requirements.operation not in operations:
|
||||
reasons.append(f"不支持操作 {requirements.operation}")
|
||||
|
||||
missing_features = sorted(set(requirements.features) - _set(capabilities.get("features")))
|
||||
if missing_features:
|
||||
reasons.append(f"缺少特性:{', '.join(missing_features)}")
|
||||
|
||||
if requirements.reference_mode and requirements.reference_mode != "none":
|
||||
if requirements.reference_mode not in _set(capabilities.get("reference_modes")):
|
||||
reasons.append(f"不支持 {requirements.reference_mode} 参考图模式")
|
||||
|
||||
_check_limit(reasons, capabilities, "max_reference_images", requirements.reference_images, "参考图片")
|
||||
_check_limit(reasons, capabilities, "max_reference_videos", requirements.reference_videos, "参考视频")
|
||||
_check_limit(reasons, capabilities, "max_reference_audios", requirements.reference_audios, "参考音频")
|
||||
|
||||
for field_name, required, label in (
|
||||
("aspect_ratios", requirements.aspect_ratio, "画面比例"),
|
||||
("resolutions", requirements.resolution, "分辨率"),
|
||||
("languages", requirements.language, "语言"),
|
||||
("output_formats", requirements.output_format, "输出格式"),
|
||||
):
|
||||
if required and required not in _set(capabilities.get(field_name)):
|
||||
reasons.append(f"不支持{label} {required}")
|
||||
|
||||
if requirements.duration is not None:
|
||||
durations = capabilities.get("durations")
|
||||
supported = set(durations) if isinstance(durations, (list, tuple, set, frozenset)) else set()
|
||||
if requirements.duration not in supported:
|
||||
reasons.append(f"不支持时长 {requirements.duration} 秒")
|
||||
|
||||
if requirements.public_voice:
|
||||
voice_map = capabilities.get("voice_map")
|
||||
if not isinstance(voice_map, dict) or not voice_map.get(requirements.public_voice):
|
||||
reasons.append(f"缺少公开音色 {requirements.public_voice} 的供应商映射")
|
||||
|
||||
if requirements.char_count is not None:
|
||||
max_chars = _positive_int(capabilities.get("max_chars"))
|
||||
if max_chars is None or requirements.char_count > max_chars:
|
||||
reasons.append(
|
||||
f"字符上限不足:需要 {requirements.char_count},配置为 {max_chars if max_chars is not None else '缺失'}"
|
||||
)
|
||||
|
||||
if requirements.speed_ratio is not None:
|
||||
speed_range = capabilities.get("speed_range")
|
||||
if not isinstance(speed_range, (list, tuple)) or len(speed_range) != 2:
|
||||
reasons.append("缺少合法的语速范围配置")
|
||||
else:
|
||||
try:
|
||||
minimum, maximum = float(speed_range[0]), float(speed_range[1])
|
||||
except (TypeError, ValueError):
|
||||
reasons.append("语速范围配置不是数字")
|
||||
else:
|
||||
if not minimum <= requirements.speed_ratio <= maximum:
|
||||
reasons.append(f"不支持语速 {requirements.speed_ratio:g}")
|
||||
|
||||
if requirements.capability == ModelConfig.Capability.VIDEO and not _video_pricing_available(
|
||||
model, requirements.resolution
|
||||
):
|
||||
reasons.append(f"缺少分辨率 {requirements.resolution} 的视频价格配置")
|
||||
|
||||
return CapabilityMatch(not reasons, tuple(reasons))
|
||||
|
||||
|
||||
def _timestamp(value: datetime | None) -> float:
|
||||
return value.timestamp() if value is not None else 0.0
|
||||
|
||||
|
||||
def _candidate_sort_key(model: ModelConfig) -> tuple[int, float, float, str]:
|
||||
return (
|
||||
provider_fallback_priority(model.provider),
|
||||
-_timestamp(model.updated_at),
|
||||
-_timestamp(model.created_at),
|
||||
str(model.id),
|
||||
)
|
||||
|
||||
|
||||
def resolve_fallback_candidates(
|
||||
*,
|
||||
primary_model: ModelConfig,
|
||||
requirements: ModelRequirements,
|
||||
attempted_model_ids: set[Any] | frozenset[Any] = frozenset(),
|
||||
excluded_provider_ids: set[Any] | frozenset[Any] = frozenset(),
|
||||
) -> list[ModelConfig]:
|
||||
"""在主模型明确失败后解析候选;不会替换或重排用户传入的主模型。"""
|
||||
|
||||
if not model_allows_fallback(primary_model):
|
||||
return []
|
||||
|
||||
attempted = {str(value) for value in attempted_model_ids}
|
||||
attempted.add(str(primary_model.id))
|
||||
excluded_providers = {str(value) for value in excluded_provider_ids}
|
||||
policy = load_model_routing_policy()
|
||||
remaining_model_slots = max(0, policy.max_models - len(attempted))
|
||||
if remaining_model_slots == 0:
|
||||
return []
|
||||
|
||||
candidates = []
|
||||
queryset = ModelConfig.objects.select_related("provider").filter(
|
||||
capability=requirements.capability,
|
||||
status=ModelConfig.Status.ACTIVE,
|
||||
provider__status=ModelProvider.Status.ACTIVE,
|
||||
)
|
||||
for model in queryset:
|
||||
if str(model.id) in attempted or str(model.provider_id) in excluded_providers:
|
||||
continue
|
||||
if not model_is_fallback_candidate(model) or capability_metadata_errors(model):
|
||||
continue
|
||||
if match_model_requirements(model, requirements).matched:
|
||||
candidates.append(model)
|
||||
|
||||
candidates.sort(key=_candidate_sort_key)
|
||||
return candidates[:remaining_model_slots]
|
||||
|
||||
|
||||
__all__ = [
|
||||
"CapabilityMatch",
|
||||
"ModelRequirements",
|
||||
"capability_metadata",
|
||||
"capability_metadata_errors",
|
||||
"match_model_requirements",
|
||||
"model_allows_fallback",
|
||||
"model_is_fallback_candidate",
|
||||
"model_metadata_errors",
|
||||
"provider_fallback_priority",
|
||||
"provider_metadata_errors",
|
||||
"resolve_fallback_candidates",
|
||||
"routing_metadata",
|
||||
]
|
||||
@@ -174,6 +174,60 @@ class AITask(TeamOwnedModel):
|
||||
return f"{self.task_type}:{self.status}:{self.id}"
|
||||
|
||||
|
||||
class AIModelAttempt(TimeStampedModel):
|
||||
"""AITask 下的一次真实模型请求审计;不参与积分预留、扣费或任务生命周期。"""
|
||||
|
||||
class Status(models.TextChoices):
|
||||
STARTED = "started", "Started"
|
||||
SUCCEEDED = "succeeded", "Succeeded"
|
||||
FAILED = "failed", "Failed"
|
||||
|
||||
task = models.ForeignKey(AITask, on_delete=models.CASCADE, related_name="model_attempts")
|
||||
sequence = models.PositiveIntegerField()
|
||||
provider = models.ForeignKey(ModelProvider, on_delete=models.SET_NULL, null=True, blank=True)
|
||||
model_config = models.ForeignKey(ModelConfig, on_delete=models.SET_NULL, null=True, blank=True)
|
||||
# 名称均保存历史快照,后续后台改名或删除配置也不影响旧任务审计。
|
||||
provider_name = models.CharField(max_length=128)
|
||||
provider_display_name = models.CharField(max_length=128, blank=True)
|
||||
model_name = models.CharField(max_length=128)
|
||||
model_display_name = models.CharField(max_length=128, blank=True)
|
||||
public_model_name = models.CharField(max_length=128, blank=True)
|
||||
capability = models.CharField(max_length=32)
|
||||
operation = models.CharField(max_length=64)
|
||||
status = models.CharField(max_length=24, choices=Status.choices, default=Status.STARTED)
|
||||
is_retry = models.BooleanField(default=False)
|
||||
is_fallback = models.BooleanField(default=False)
|
||||
previous_attempt = models.ForeignKey(
|
||||
"self",
|
||||
on_delete=models.SET_NULL,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="next_attempts",
|
||||
)
|
||||
provider_task_id = models.CharField(max_length=255, blank=True)
|
||||
started_at = models.DateTimeField()
|
||||
finished_at = models.DateTimeField(null=True, blank=True)
|
||||
duration_ms = models.PositiveBigIntegerField(null=True, blank=True)
|
||||
error_type = models.CharField(max_length=64, blank=True)
|
||||
provider_error_code = models.CharField(max_length=128, blank=True)
|
||||
raw_error = models.TextField(blank=True)
|
||||
safe_error_summary = models.TextField(blank=True)
|
||||
usage = models.JSONField(default=dict, blank=True)
|
||||
platform_cost = models.DecimalField(max_digits=12, decimal_places=4, default=0)
|
||||
request_summary = models.JSONField(default=dict, blank=True)
|
||||
response_summary = models.JSONField(default=dict, blank=True)
|
||||
|
||||
class Meta:
|
||||
ordering = ["sequence"]
|
||||
constraints = [
|
||||
models.UniqueConstraint(fields=["task", "sequence"], name="ai_attempt_task_sequence_unique"),
|
||||
]
|
||||
indexes = [models.Index(fields=["task", "status"], name="ai_attempt_task_status_idx")]
|
||||
|
||||
def __str__(self) -> str:
|
||||
return f"{self.task_id}:{self.sequence}:{self.status}"
|
||||
|
||||
|
||||
class QualityWord(TimeStampedModel):
|
||||
"""平台单层质量词配置:按生成阶段(stage)+ 槽位(slot)挂若干质量词,
|
||||
生成侧拼提示词时优先读取本配置;**无配置则回落各 builder 的写死值**(零回归)。
|
||||
|
||||
@@ -86,6 +86,7 @@ class OpenAICompatibleProvider(VolcanoArkProvider):
|
||||
endpoint: str = "images/generations",
|
||||
image: str | list[str] | None = None,
|
||||
size: str = "1024x1536",
|
||||
timeout: float = 300,
|
||||
) -> dict[str, Any]:
|
||||
"""文生图(可选单图参考 base64)。多图参考请用 image_edit。返回体含 url 或 b64_json。"""
|
||||
if not self.api_key:
|
||||
@@ -98,7 +99,7 @@ class OpenAICompatibleProvider(VolcanoArkProvider):
|
||||
self._endpoint_url(endpoint),
|
||||
headers={"Authorization": f"Bearer {self.api_key}", "Content-Type": "application/json"},
|
||||
json=body,
|
||||
timeout=300,
|
||||
timeout=timeout,
|
||||
)
|
||||
_raise_with_body(response)
|
||||
return response.json()
|
||||
@@ -111,6 +112,7 @@ class OpenAICompatibleProvider(VolcanoArkProvider):
|
||||
images: list[str],
|
||||
endpoint: str = "images/edits",
|
||||
size: str = "1024x1536",
|
||||
timeout: float = 300,
|
||||
) -> dict[str, Any]:
|
||||
"""参考图编辑/合成(gpt-image-2 核心能力):multipart `image[]` 上传一张或多张参考图。
|
||||
|
||||
@@ -139,7 +141,47 @@ class OpenAICompatibleProvider(VolcanoArkProvider):
|
||||
headers={"Authorization": f"Bearer {self.api_key}"}, # multipart 不要手设 Content-Type
|
||||
files=files,
|
||||
data=data,
|
||||
timeout=300,
|
||||
timeout=timeout,
|
||||
)
|
||||
_raise_with_body(response)
|
||||
return response.json()
|
||||
|
||||
def synthesize(
|
||||
self,
|
||||
*,
|
||||
model: str,
|
||||
text: str,
|
||||
voice_type: str,
|
||||
speed_ratio: float = 1.0,
|
||||
endpoint: str = "audio/speech",
|
||||
output_format: str = "mp3",
|
||||
timeout: float = 60,
|
||||
uid: str = "airshelf",
|
||||
) -> tuple[bytes, int]:
|
||||
"""OpenAI 兼容 ``audio/speech``:返回音频字节和可选时长毫秒。
|
||||
|
||||
``uid`` 仅为与火山直连调用签名一致,不发送给标准 OpenAI 请求。新增供应商只需在
|
||||
ModelConfig 配 endpoint / voice_map,无需新增 Provider 分支。
|
||||
"""
|
||||
del uid
|
||||
if not self.api_key:
|
||||
raise ValueError("中转站 api_key 未配置")
|
||||
response = requests.post(
|
||||
self._endpoint_url(endpoint or "audio/speech"),
|
||||
headers={"Authorization": f"Bearer {self.api_key}", "Content-Type": "application/json"},
|
||||
json={
|
||||
"model": model,
|
||||
"input": text,
|
||||
"voice": voice_type,
|
||||
"speed": float(speed_ratio or 1.0),
|
||||
"response_format": output_format,
|
||||
},
|
||||
timeout=timeout,
|
||||
)
|
||||
_raise_with_body(response)
|
||||
duration_ms = 0
|
||||
try:
|
||||
duration_ms = int(float(response.headers.get("x-audio-duration-ms") or 0))
|
||||
except (TypeError, ValueError):
|
||||
duration_ms = 0
|
||||
return response.content, duration_ms
|
||||
|
||||
@@ -59,7 +59,14 @@ class VolcanoArkProvider:
|
||||
payload=data,
|
||||
)
|
||||
|
||||
def chat_completion(self, *, model: str, messages: list[dict[str, str]], endpoint: str = "chat/completions") -> dict[str, Any]:
|
||||
def chat_completion(
|
||||
self,
|
||||
*,
|
||||
model: str,
|
||||
messages: list[dict[str, str]],
|
||||
endpoint: str = "chat/completions",
|
||||
timeout: float = 120,
|
||||
) -> dict[str, Any]:
|
||||
if not self.api_key:
|
||||
raise ValueError("VOLCANO_ARK_API_KEY is not configured")
|
||||
|
||||
@@ -67,7 +74,7 @@ class VolcanoArkProvider:
|
||||
f"{self.base_url.rstrip('/')}/{endpoint.lstrip('/')}",
|
||||
headers={"Authorization": f"Bearer {self.api_key}", "Content-Type": "application/json"},
|
||||
json={"model": model, "messages": messages},
|
||||
timeout=120,
|
||||
timeout=timeout,
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
@@ -80,6 +87,7 @@ class VolcanoArkProvider:
|
||||
endpoint: str = "chat/completions",
|
||||
temperature: float = 0.8,
|
||||
extra_body: dict[str, Any] | None = None,
|
||||
timeout: float = 300,
|
||||
) -> Iterator[dict[str, Any]]:
|
||||
"""流式对话:逐块 yield {type:'delta'|'tool_call'|'done', ...}。
|
||||
OpenAI 兼容 SSE(火山 ARK / 各中转站同构),供脚本 agent 的 SSE 端实时转发。"""
|
||||
@@ -97,7 +105,7 @@ class VolcanoArkProvider:
|
||||
},
|
||||
json=body,
|
||||
stream=True,
|
||||
timeout=300,
|
||||
timeout=timeout,
|
||||
) as response:
|
||||
response.raise_for_status()
|
||||
# SSE 响应常不带 charset,requests 会按 latin-1 解码 → 中文乱码。强制 UTF-8。
|
||||
@@ -157,6 +165,7 @@ class VolcanoArkProvider:
|
||||
endpoint: str = "images/generations",
|
||||
image: str | list[str] | None = None,
|
||||
size: str = "2K",
|
||||
timeout: float = 180,
|
||||
) -> dict[str, Any]:
|
||||
if not self.api_key:
|
||||
raise ValueError("VOLCANO_ARK_API_KEY is not configured")
|
||||
@@ -174,7 +183,7 @@ class VolcanoArkProvider:
|
||||
f"{self.base_url.rstrip('/')}/{endpoint.lstrip('/')}",
|
||||
headers={"Authorization": f"Bearer {self.api_key}", "Content-Type": "application/json"},
|
||||
json=body,
|
||||
timeout=180,
|
||||
timeout=timeout,
|
||||
)
|
||||
# 非 2xx 抽出火山真实报错(error.code/message)而非裸 HTTP 400 —— 否则 Seedream 的
|
||||
# 内容审核 / 尺寸过小(历史 4:5 套图 raise_for_status 吞 body 的踩坑)等真因看不见。
|
||||
@@ -195,6 +204,7 @@ class VolcanoArkProvider:
|
||||
content_items: list[dict[str, Any]] | None = None,
|
||||
seed: int | None = None,
|
||||
search_mode: str = "off",
|
||||
timeout: float = 120,
|
||||
) -> dict[str, Any]:
|
||||
if not self.api_key:
|
||||
raise ValueError("VOLCANO_ARK_API_KEY is not configured")
|
||||
@@ -223,18 +233,24 @@ class VolcanoArkProvider:
|
||||
f"{self.base_url.rstrip('/')}/{endpoint.lstrip('/')}",
|
||||
headers={"Authorization": f"Bearer {self.api_key}", "Content-Type": "application/json"},
|
||||
json=body,
|
||||
timeout=120,
|
||||
timeout=timeout,
|
||||
)
|
||||
_raise_volcano_error(response)
|
||||
return response.json()
|
||||
|
||||
def poll_video_task(self, *, endpoint: str, provider_task_id: str) -> dict[str, Any]:
|
||||
def poll_video_task(
|
||||
self,
|
||||
*,
|
||||
endpoint: str,
|
||||
provider_task_id: str,
|
||||
timeout: float = 60,
|
||||
) -> dict[str, Any]:
|
||||
if not self.api_key:
|
||||
raise ValueError("VOLCANO_ARK_API_KEY is not configured")
|
||||
response = requests.get(
|
||||
f"{self.base_url.rstrip('/')}/{endpoint.rstrip('/')}/{provider_task_id}",
|
||||
headers={"Authorization": f"Bearer {self.api_key}"},
|
||||
timeout=60,
|
||||
timeout=timeout,
|
||||
)
|
||||
_raise_volcano_error(response)
|
||||
return response.json()
|
||||
@@ -305,7 +321,15 @@ class VolcanoTtsProvider:
|
||||
def configured(self) -> bool:
|
||||
return bool(self.appid and self.access_token)
|
||||
|
||||
def synthesize(self, *, text: str, voice_type: str, speed_ratio: float = 1.0, uid: str = "airshelf") -> tuple[bytes, int]:
|
||||
def synthesize(
|
||||
self,
|
||||
*,
|
||||
text: str,
|
||||
voice_type: str,
|
||||
speed_ratio: float = 1.0,
|
||||
uid: str = "airshelf",
|
||||
timeout: float = 60,
|
||||
) -> tuple[bytes, int]:
|
||||
"""合成一段语音。返回 (mp3 字节, 时长毫秒;接口没回时长则为 0)。"""
|
||||
if not self.configured:
|
||||
raise TtsNotConfigured(
|
||||
@@ -322,7 +346,7 @@ class VolcanoTtsProvider:
|
||||
self.base_url,
|
||||
headers={"Authorization": f"Bearer;{self.access_token}"},
|
||||
json=body,
|
||||
timeout=60,
|
||||
timeout=timeout,
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
|
||||
@@ -26,6 +26,7 @@ class YunqiProvider(VolcanoArkProvider):
|
||||
endpoint: str = "images/generations",
|
||||
image: str | list[str] | None = None,
|
||||
size: str = "1024x1536",
|
||||
timeout: float = 300,
|
||||
) -> dict[str, Any]:
|
||||
if not self.api_key:
|
||||
raise ValueError("YUNQI_API_KEY is not configured")
|
||||
@@ -37,7 +38,7 @@ class YunqiProvider(VolcanoArkProvider):
|
||||
f"{self.base_url.rstrip('/')}/{endpoint.lstrip('/')}",
|
||||
headers={"Authorization": f"Bearer {self.api_key}", "Content-Type": "application/json"},
|
||||
json=body,
|
||||
timeout=300,
|
||||
timeout=timeout,
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
@@ -0,0 +1,456 @@
|
||||
"""统一模型请求执行器:重试、动态切换和每次真实调用的 Feedback 留痕。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from decimal import Decimal, InvalidOperation
|
||||
import random
|
||||
import re
|
||||
import time
|
||||
from typing import Any, Callable, Generic, Mapping, TypeVar
|
||||
|
||||
from django.db import transaction
|
||||
from django.db.models import F, Max
|
||||
from django.utils import timezone
|
||||
|
||||
from apps.ai.generation_errors import PublicGenerationError, classify_generation_error
|
||||
from apps.ai.model_routing import ModelRequirements, resolve_fallback_candidates
|
||||
from apps.ai.models import AIModelAttempt, AITask, ModelConfig
|
||||
from apps.ai.routing_policy import load_model_routing_policy
|
||||
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AttemptMetadata:
|
||||
"""调用方可返回的审计摘要;不得放入密钥、鉴权头或完整素材。"""
|
||||
|
||||
provider_task_id: str = ""
|
||||
usage: Mapping[str, Any] = field(default_factory=dict)
|
||||
platform_cost: Decimal = Decimal("0")
|
||||
response_summary: Mapping[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RoutingExecutionResult(Generic[T]):
|
||||
value: T
|
||||
actual_model: ModelConfig
|
||||
attempt: AIModelAttempt
|
||||
call_count: int
|
||||
fallback_used: bool
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ErrorRoutingDecision:
|
||||
retry_current: bool
|
||||
fallback: bool
|
||||
exclude_provider: bool = False
|
||||
|
||||
|
||||
_TERMINAL_ERRORS = {
|
||||
"content_rejected",
|
||||
"invalid_input",
|
||||
"asset_unavailable",
|
||||
"user_credit_insufficient",
|
||||
"processing_failed",
|
||||
}
|
||||
_FALLBACK_ONLY_ERRORS = {"provider_quota_exhausted", "provider_config_error", "model_unavailable"}
|
||||
_PROVIDER_SCOPE_ERRORS = {"provider_quota_exhausted", "provider_config_error"}
|
||||
_SENSITIVE_KEYS = {"api_key", "apikey", "authorization", "password", "secret", "access_token", "refresh_token"}
|
||||
_SECRET_PATTERNS = (
|
||||
re.compile(r"(?i)(authorization\s*[:=]\s*)([^\s,;]+)"),
|
||||
re.compile(r"(?i)((?:api[_-]?key|access[_-]?token|secret)\s*[:=]\s*)([^\s,;]+)"),
|
||||
re.compile(r"\bsk-[A-Za-z0-9_-]{8,}\b"),
|
||||
)
|
||||
|
||||
|
||||
def decide_error_routing(error: PublicGenerationError) -> ErrorRoutingDecision:
|
||||
"""把统一安全错误映射成路由动作;业务入口不得复制错误字符串判断。"""
|
||||
|
||||
if error.code in _TERMINAL_ERRORS:
|
||||
return ErrorRoutingDecision(retry_current=False, fallback=False)
|
||||
if error.code in _FALLBACK_ONLY_ERRORS:
|
||||
return ErrorRoutingDecision(
|
||||
retry_current=False,
|
||||
fallback=True,
|
||||
exclude_provider=error.code in _PROVIDER_SCOPE_ERRORS,
|
||||
)
|
||||
# 网络、超时、限流、5xx 及暂未精确分类的 Provider 异常:主模型按能力策略重试后再切换。
|
||||
return ErrorRoutingDecision(retry_current=True, fallback=True)
|
||||
|
||||
|
||||
def _redact_text(value: Any, limit: int = 4000) -> str:
|
||||
text = str(value or "")
|
||||
for pattern in _SECRET_PATTERNS:
|
||||
if pattern.groups >= 2:
|
||||
text = pattern.sub(r"\1[REDACTED]", text)
|
||||
else:
|
||||
text = pattern.sub("[REDACTED]", text)
|
||||
return text[:limit]
|
||||
|
||||
|
||||
def sanitize_attempt_summary(value: Any, *, depth: int = 0) -> Any:
|
||||
"""生成可持久化的轻量脱敏摘要,避免日志保存密钥、二进制或完整长素材。"""
|
||||
|
||||
if depth >= 4:
|
||||
return "[摘要层级已截断]"
|
||||
if isinstance(value, Mapping):
|
||||
result = {}
|
||||
for index, (key, item) in enumerate(value.items()):
|
||||
if index >= 40:
|
||||
result["_truncated"] = True
|
||||
break
|
||||
name = str(key)
|
||||
if name.lower() in _SENSITIVE_KEYS:
|
||||
result[name] = "[REDACTED]"
|
||||
else:
|
||||
result[name] = sanitize_attempt_summary(item, depth=depth + 1)
|
||||
return result
|
||||
if isinstance(value, (list, tuple, set, frozenset)):
|
||||
items = list(value)
|
||||
return [sanitize_attempt_summary(item, depth=depth + 1) for item in items[:20]]
|
||||
if isinstance(value, (bytes, bytearray, memoryview)):
|
||||
return f"[二进制 {len(value)} 字节]"
|
||||
if value is None or isinstance(value, (bool, int, float)):
|
||||
return value
|
||||
if isinstance(value, Decimal):
|
||||
return str(value)
|
||||
return _redact_text(value, limit=500)
|
||||
|
||||
|
||||
def _coerce_metadata(value: Any) -> AttemptMetadata:
|
||||
if isinstance(value, AttemptMetadata):
|
||||
return AttemptMetadata(
|
||||
provider_task_id=str(value.provider_task_id or "")[:255],
|
||||
usage=value.usage if isinstance(value.usage, Mapping) else {},
|
||||
platform_cost=value.platform_cost,
|
||||
response_summary=value.response_summary if isinstance(value.response_summary, Mapping) else {},
|
||||
)
|
||||
if not isinstance(value, Mapping):
|
||||
return AttemptMetadata()
|
||||
raw_cost = value.get("platform_cost", 0)
|
||||
try:
|
||||
cost = Decimal(str(raw_cost or 0))
|
||||
except (InvalidOperation, TypeError, ValueError):
|
||||
cost = Decimal("0")
|
||||
return AttemptMetadata(
|
||||
provider_task_id=str(value.get("provider_task_id") or value.get("task_id") or "")[:255],
|
||||
usage=value.get("usage") if isinstance(value.get("usage"), Mapping) else {},
|
||||
platform_cost=max(cost, Decimal("0")),
|
||||
response_summary=value.get("response_summary")
|
||||
if isinstance(value.get("response_summary"), Mapping)
|
||||
else {},
|
||||
)
|
||||
|
||||
|
||||
def _default_result_metadata(result: Any, model: ModelConfig) -> AttemptMetadata:
|
||||
del model
|
||||
if not isinstance(result, Mapping):
|
||||
return AttemptMetadata()
|
||||
provider_task_id = result.get("provider_task_id") or result.get("task_id") or result.get("id") or ""
|
||||
usage = result.get("usage") if isinstance(result.get("usage"), Mapping) else {}
|
||||
return AttemptMetadata(
|
||||
provider_task_id=str(provider_task_id)[:255],
|
||||
usage=usage,
|
||||
response_summary={"provider_task_id": str(provider_task_id)[:255]} if provider_task_id else {},
|
||||
)
|
||||
|
||||
|
||||
def _default_error_metadata(exc: Exception, model: ModelConfig) -> AttemptMetadata:
|
||||
del model
|
||||
return _coerce_metadata(getattr(exc, "attempt_metadata", None))
|
||||
|
||||
|
||||
def _provider_error_code(exc: Exception) -> str:
|
||||
response = getattr(exc, "response", None)
|
||||
if response is not None:
|
||||
try:
|
||||
payload = response.json()
|
||||
error = payload.get("error") if isinstance(payload, Mapping) else None
|
||||
if isinstance(error, Mapping):
|
||||
return str(error.get("code") or error.get("type") or "")[:128]
|
||||
except Exception: # noqa: BLE001 - Provider 错误响应不保证 JSON 格式。
|
||||
pass
|
||||
return str(getattr(exc, "code", "") or "")[:128]
|
||||
|
||||
|
||||
@transaction.atomic
|
||||
def _start_attempt(
|
||||
*,
|
||||
task: AITask,
|
||||
model: ModelConfig,
|
||||
requirements: ModelRequirements,
|
||||
public_model_name: str,
|
||||
is_retry: bool,
|
||||
is_fallback: bool,
|
||||
previous_attempt: AIModelAttempt | None,
|
||||
request_summary: Mapping[str, Any],
|
||||
) -> AIModelAttempt:
|
||||
# 锁父任务只用于分配稳定序号,不修改父任务,因而不会刷新 AITask.updated_at。
|
||||
AITask.objects.select_for_update().only("id").get(pk=task.pk)
|
||||
sequence = (
|
||||
AIModelAttempt.objects.filter(task_id=task.pk).aggregate(max_sequence=Max("sequence"))["max_sequence"] or 0
|
||||
) + 1
|
||||
provider = model.provider
|
||||
return AIModelAttempt.objects.create(
|
||||
task_id=task.pk,
|
||||
sequence=sequence,
|
||||
provider=provider,
|
||||
model_config=model,
|
||||
provider_name=provider.name,
|
||||
provider_display_name=provider.display_name,
|
||||
model_name=model.name,
|
||||
model_display_name=model.display_name,
|
||||
public_model_name=public_model_name,
|
||||
capability=requirements.capability,
|
||||
operation=requirements.operation,
|
||||
is_retry=is_retry,
|
||||
is_fallback=is_fallback,
|
||||
previous_attempt=previous_attempt,
|
||||
started_at=timezone.now(),
|
||||
request_summary=sanitize_attempt_summary(request_summary),
|
||||
)
|
||||
|
||||
|
||||
def _finish_attempt(
|
||||
attempt: AIModelAttempt,
|
||||
*,
|
||||
status: str,
|
||||
elapsed_seconds: float,
|
||||
metadata: AttemptMetadata,
|
||||
error: PublicGenerationError | None = None,
|
||||
exc: Exception | None = None,
|
||||
) -> None:
|
||||
cost = max(Decimal(metadata.platform_cost), Decimal("0")).quantize(Decimal("0.0001"))
|
||||
attempt.status = status
|
||||
attempt.finished_at = timezone.now()
|
||||
attempt.duration_ms = max(0, round(elapsed_seconds * 1000))
|
||||
attempt.provider_task_id = metadata.provider_task_id[:255]
|
||||
attempt.usage = sanitize_attempt_summary(metadata.usage)
|
||||
attempt.platform_cost = cost
|
||||
attempt.response_summary = sanitize_attempt_summary(metadata.response_summary)
|
||||
if error is not None:
|
||||
attempt.error_type = error.code
|
||||
attempt.provider_error_code = _provider_error_code(exc) if exc is not None else ""
|
||||
attempt.raw_error = _redact_text(exc)
|
||||
attempt.safe_error_summary = error.fallback_message
|
||||
attempt.save(
|
||||
update_fields=[
|
||||
"status",
|
||||
"finished_at",
|
||||
"duration_ms",
|
||||
"provider_task_id",
|
||||
"usage",
|
||||
"platform_cost",
|
||||
"response_summary",
|
||||
"error_type",
|
||||
"provider_error_code",
|
||||
"raw_error",
|
||||
"safe_error_summary",
|
||||
"updated_at",
|
||||
]
|
||||
)
|
||||
if cost > 0:
|
||||
# QuerySet.update 不触发父任务 auto_now,避免尝试日志影响僵尸任务回收时间口径。
|
||||
AITask.objects.filter(pk=attempt.task_id).update(base_cost=F("base_cost") + cost)
|
||||
|
||||
|
||||
def _retry_after_seconds(exc: Exception) -> float | None:
|
||||
response = getattr(exc, "response", None)
|
||||
headers = getattr(response, "headers", None)
|
||||
if not isinstance(headers, Mapping):
|
||||
return None
|
||||
raw = headers.get("Retry-After") or headers.get("retry-after")
|
||||
try:
|
||||
value = float(raw)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
return max(0.0, value)
|
||||
|
||||
|
||||
def _ability_policy(requirements: ModelRequirements):
|
||||
policy = load_model_routing_policy()
|
||||
if requirements.capability == ModelConfig.Capability.TEXT:
|
||||
request_timeout = policy.text.stream_timeout if "streaming" in requirements.features else policy.text.request_timeout
|
||||
return policy, policy.text.retry_delays, request_timeout, policy.text.total_timeout, policy.text.retry_after_cap
|
||||
if requirements.capability == ModelConfig.Capability.IMAGE:
|
||||
return policy, policy.image.retry_delays, policy.image.request_timeout, policy.image.total_timeout, policy.image.retry_after_cap
|
||||
if requirements.capability == ModelConfig.Capability.AUDIO:
|
||||
return policy, policy.audio.retry_delays, policy.audio.request_timeout, policy.audio.total_timeout, policy.audio.retry_after_cap
|
||||
if requirements.capability == ModelConfig.Capability.VIDEO:
|
||||
return (
|
||||
policy,
|
||||
policy.video.submit_retry_delays,
|
||||
policy.video.submit_timeout,
|
||||
policy.video.submit_total_timeout,
|
||||
policy.video.retry_after_cap,
|
||||
)
|
||||
raise ValueError(f"暂不支持能力 {requirements.capability} 的统一模型路由")
|
||||
|
||||
|
||||
def execute_model_call(
|
||||
*,
|
||||
task: AITask,
|
||||
primary_model: ModelConfig,
|
||||
requirements: ModelRequirements,
|
||||
public_model_name: str,
|
||||
invoke: Callable[[ModelConfig, float], T],
|
||||
request_summary: Mapping[str, Any] | Callable[[ModelConfig], Mapping[str, Any]] | None = None,
|
||||
result_metadata: Callable[[T, ModelConfig], AttemptMetadata | Mapping[str, Any]] | None = None,
|
||||
error_metadata: Callable[[Exception, ModelConfig], AttemptMetadata | Mapping[str, Any]] | None = None,
|
||||
error_classifier: Callable[[Exception, ModelConfig], PublicGenerationError] | None = None,
|
||||
candidate_resolver: Callable[..., list[ModelConfig]] | None = None,
|
||||
sleep: Callable[[float], None] = time.sleep,
|
||||
monotonic: Callable[[], float] = time.monotonic,
|
||||
uniform: Callable[[float, float], float] = random.uniform,
|
||||
) -> RoutingExecutionResult[T]:
|
||||
"""在一个既有 AITask 内完成真实调用;不创建预留、不扣费、不释放积分。"""
|
||||
|
||||
if task.model_config_id != primary_model.id:
|
||||
raise ValueError("AITask.model_config 必须保持为用户选择或系统默认的主模型")
|
||||
if primary_model.capability != requirements.capability:
|
||||
raise ValueError("主模型 capability 与本次 ModelRequirements 不一致")
|
||||
|
||||
policy, retry_delays, request_timeout, total_timeout, retry_after_cap = _ability_policy(requirements)
|
||||
resolver = candidate_resolver or resolve_fallback_candidates
|
||||
result_meta = result_metadata or _default_result_metadata
|
||||
error_meta = error_metadata or _default_error_metadata
|
||||
classify = error_classifier or (
|
||||
lambda exc, model: classify_generation_error(
|
||||
exc,
|
||||
operation=requirements.operation,
|
||||
provider_name=model.provider.name,
|
||||
reference_id=str(task.id),
|
||||
)
|
||||
)
|
||||
|
||||
route_started = monotonic()
|
||||
current_model = primary_model
|
||||
attempted_model_ids: set[Any] = set()
|
||||
excluded_provider_ids: set[Any] = set()
|
||||
previous_attempt: AIModelAttempt | None = None
|
||||
calls = 0
|
||||
primary_retry_index = 0
|
||||
fallback_used = False
|
||||
last_exception: Exception | None = None
|
||||
|
||||
while calls < policy.max_calls:
|
||||
remaining = total_timeout - (monotonic() - route_started)
|
||||
if remaining <= 0:
|
||||
if last_exception is not None:
|
||||
raise last_exception
|
||||
raise TimeoutError("模型路由总时限已耗尽")
|
||||
|
||||
is_retry = current_model.id == primary_model.id and primary_retry_index > 0
|
||||
summary = request_summary(current_model) if callable(request_summary) else (request_summary or {})
|
||||
attempt = _start_attempt(
|
||||
task=task,
|
||||
model=current_model,
|
||||
requirements=requirements,
|
||||
public_model_name=public_model_name,
|
||||
is_retry=is_retry,
|
||||
is_fallback=fallback_used,
|
||||
previous_attempt=previous_attempt,
|
||||
request_summary=summary,
|
||||
)
|
||||
calls += 1
|
||||
call_started = monotonic()
|
||||
try:
|
||||
value = invoke(current_model, min(float(request_timeout), max(0.001, remaining)))
|
||||
except Exception as exc:
|
||||
elapsed = monotonic() - call_started
|
||||
public_error = classify(exc, current_model)
|
||||
decision = decide_error_routing(public_error)
|
||||
try:
|
||||
failed_metadata = _coerce_metadata(error_meta(exc, current_model))
|
||||
except Exception as metadata_exc: # noqa: BLE001 - 审计摘要失败不能遮蔽原始 Provider 异常。
|
||||
failed_metadata = AttemptMetadata(
|
||||
response_summary={"metadata_error": _redact_text(metadata_exc, limit=300)}
|
||||
)
|
||||
# 异步视频若已经拿到 Provider 任务 ID,表示提交结果并非“明确未创建”;禁止盲目重提。
|
||||
if requirements.capability == ModelConfig.Capability.VIDEO and failed_metadata.provider_task_id:
|
||||
decision = ErrorRoutingDecision(retry_current=False, fallback=False)
|
||||
_finish_attempt(
|
||||
attempt,
|
||||
status=AIModelAttempt.Status.FAILED,
|
||||
elapsed_seconds=elapsed,
|
||||
metadata=failed_metadata,
|
||||
error=public_error,
|
||||
exc=exc,
|
||||
)
|
||||
previous_attempt = attempt
|
||||
last_exception = exc
|
||||
if decision.exclude_provider:
|
||||
excluded_provider_ids.add(current_model.provider_id)
|
||||
if not decision.fallback:
|
||||
raise
|
||||
|
||||
can_retry_primary = (
|
||||
not fallback_used
|
||||
and decision.retry_current
|
||||
and primary_retry_index < len(retry_delays)
|
||||
and calls < policy.max_calls
|
||||
)
|
||||
if can_retry_primary:
|
||||
retry_after = _retry_after_seconds(exc)
|
||||
if retry_after is not None and retry_after > retry_after_cap:
|
||||
can_retry_primary = False
|
||||
else:
|
||||
base_delay = retry_after if retry_after is not None else float(retry_delays[primary_retry_index])
|
||||
delay = base_delay if retry_after is not None else base_delay * (1 + uniform(-policy.jitter_ratio, policy.jitter_ratio))
|
||||
if monotonic() - route_started + max(0.0, delay) >= total_timeout:
|
||||
can_retry_primary = False
|
||||
else:
|
||||
sleep(max(0.0, delay))
|
||||
primary_retry_index += 1
|
||||
if can_retry_primary:
|
||||
continue
|
||||
|
||||
attempted_model_ids.add(current_model.id)
|
||||
if calls >= policy.max_calls:
|
||||
raise
|
||||
candidates = resolver(
|
||||
primary_model=current_model,
|
||||
requirements=requirements,
|
||||
attempted_model_ids=attempted_model_ids,
|
||||
excluded_provider_ids=excluded_provider_ids,
|
||||
)
|
||||
if not candidates:
|
||||
raise
|
||||
current_model = candidates[0]
|
||||
fallback_used = True
|
||||
primary_retry_index = 0
|
||||
continue
|
||||
|
||||
try:
|
||||
metadata = _coerce_metadata(result_meta(value, current_model))
|
||||
except Exception as metadata_exc: # noqa: BLE001 - 内容已生成,不能因审计摘要失败重新调用模型。
|
||||
metadata = AttemptMetadata(response_summary={"metadata_error": _redact_text(metadata_exc, limit=300)})
|
||||
_finish_attempt(
|
||||
attempt,
|
||||
status=AIModelAttempt.Status.SUCCEEDED,
|
||||
elapsed_seconds=monotonic() - call_started,
|
||||
metadata=metadata,
|
||||
)
|
||||
return RoutingExecutionResult(
|
||||
value=value,
|
||||
actual_model=current_model,
|
||||
attempt=attempt,
|
||||
call_count=calls,
|
||||
fallback_used=fallback_used,
|
||||
)
|
||||
|
||||
if last_exception is not None:
|
||||
raise last_exception
|
||||
raise TimeoutError("模型路由调用次数已耗尽")
|
||||
|
||||
|
||||
__all__ = [
|
||||
"AttemptMetadata",
|
||||
"ErrorRoutingDecision",
|
||||
"RoutingExecutionResult",
|
||||
"decide_error_routing",
|
||||
"execute_model_call",
|
||||
"sanitize_attempt_summary",
|
||||
]
|
||||
@@ -0,0 +1,303 @@
|
||||
"""模型路由策略的集中读取与校验。
|
||||
|
||||
本模块只负责把 Django settings 中的 ``MODEL_ROUTING_POLICY`` 转成不可变配置对象,
|
||||
不执行重试、Fallback、Provider 调用或账务操作。业务入口后续只消费这里返回的策略,
|
||||
不得各自维护 timeout / sleep / attempts 常量。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from django.conf import settings
|
||||
|
||||
|
||||
class RoutingPolicyConfigurationError(ValueError):
|
||||
"""模型路由策略格式或取值不合法。"""
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TextRoutingPolicy:
|
||||
retry_delays: tuple[float, ...]
|
||||
request_timeout: float
|
||||
stream_timeout: float
|
||||
total_timeout: float
|
||||
retry_after_cap: float
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RequestRoutingPolicy:
|
||||
retry_delays: tuple[float, ...]
|
||||
request_timeout: float
|
||||
total_timeout: float
|
||||
retry_after_cap: float
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class VideoRoutingPolicy:
|
||||
submit_retry_delays: tuple[float, ...]
|
||||
submit_timeout: float
|
||||
submit_total_timeout: float
|
||||
poll_request_timeout: float
|
||||
generation_timeout: float
|
||||
retry_after_cap: float
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PostprocessRoutingPolicy:
|
||||
retry_delays: tuple[float, ...]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ModelRoutingPolicy:
|
||||
max_models: int
|
||||
max_calls: int
|
||||
jitter_ratio: float
|
||||
text: TextRoutingPolicy
|
||||
image: RequestRoutingPolicy
|
||||
audio: RequestRoutingPolicy
|
||||
video: VideoRoutingPolicy
|
||||
postprocess: PostprocessRoutingPolicy
|
||||
|
||||
|
||||
_LABELS = {
|
||||
"max_models": "最多尝试模型数量",
|
||||
"max_calls": "最多真实模型调用次数",
|
||||
"jitter_ratio": "重试随机抖动比例",
|
||||
"retry_delays": "重试等待时间列表",
|
||||
"request_timeout": "单次调用超时",
|
||||
"stream_timeout": "单次流式调用超时",
|
||||
"total_timeout": "逻辑任务总时限",
|
||||
"retry_after_cap": "Retry-After 最长等待时间",
|
||||
"submit_retry_delays": "视频提交重试等待时间列表",
|
||||
"submit_timeout": "视频单次提交超时",
|
||||
"submit_total_timeout": "视频提交阶段总时限",
|
||||
"poll_request_timeout": "视频单次轮询超时",
|
||||
"generation_timeout": "视频成片等待总时限",
|
||||
}
|
||||
|
||||
|
||||
def _error(path: str, message: str) -> RoutingPolicyConfigurationError:
|
||||
key = path.rsplit(".", 1)[-1].split("[", 1)[0]
|
||||
label = _LABELS.get(key, path)
|
||||
return RoutingPolicyConfigurationError(f"{path}({label}){message}")
|
||||
|
||||
|
||||
def _mapping(value: Any, path: str) -> Mapping[str, Any]:
|
||||
if not isinstance(value, Mapping):
|
||||
raise _error(path, f"必须是对象,当前值:{value!r}")
|
||||
return value
|
||||
|
||||
|
||||
def _required(source: Mapping[str, Any], key: str, path: str) -> Any:
|
||||
if key not in source:
|
||||
raise _error(f"{path}.{key}", "为必填配置")
|
||||
return source[key]
|
||||
|
||||
|
||||
def _integer(value: Any, path: str, *, minimum: int, maximum: int) -> int:
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
raise _error(path, f"必须是整数,当前值:{value!r}")
|
||||
if not minimum <= value <= maximum:
|
||||
raise _error(path, f"必须在 {minimum}~{maximum} 之间,当前值:{value!r}")
|
||||
return value
|
||||
|
||||
|
||||
def _number(value: Any, path: str, *, minimum: float, maximum: float) -> float:
|
||||
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
||||
raise _error(path, f"必须是数字,当前值:{value!r}")
|
||||
result = float(value)
|
||||
if not minimum <= result <= maximum:
|
||||
raise _error(path, f"必须在 {minimum:g}~{maximum:g} 之间,当前值:{value!r}")
|
||||
return result
|
||||
|
||||
|
||||
def _delays(value: Any, path: str) -> tuple[float, ...]:
|
||||
if isinstance(value, (str, bytes)) or not isinstance(value, Sequence):
|
||||
raise _error(path, f"必须是数字列表,当前值:{value!r}")
|
||||
return tuple(
|
||||
_number(item, f"{path}[{index}]", minimum=0, maximum=3600)
|
||||
for index, item in enumerate(value)
|
||||
)
|
||||
|
||||
|
||||
def _ensure_not_greater(*, smaller: float, larger: float, smaller_path: str, larger_path: str) -> None:
|
||||
if smaller > larger:
|
||||
raise _error(
|
||||
smaller_path,
|
||||
f"不得大于 {larger_path}(任务总时限),当前为 {smaller:g} > {larger:g}",
|
||||
)
|
||||
|
||||
|
||||
def _field_number(
|
||||
source: Mapping[str, Any],
|
||||
key: str,
|
||||
path: str,
|
||||
*,
|
||||
minimum: float = 0,
|
||||
maximum: float = 86400,
|
||||
) -> float:
|
||||
return _number(_required(source, key, path), f"{path}.{key}", minimum=minimum, maximum=maximum)
|
||||
|
||||
|
||||
def _field_delays(source: Mapping[str, Any], key: str, path: str) -> tuple[float, ...]:
|
||||
return _delays(_required(source, key, path), f"{path}.{key}")
|
||||
|
||||
|
||||
def _validate_budget(
|
||||
*,
|
||||
total: float,
|
||||
total_path: str,
|
||||
bounded_values: Mapping[str, float],
|
||||
delays: tuple[float, ...] = (),
|
||||
delays_path: str = "",
|
||||
) -> None:
|
||||
for value_path, value in bounded_values.items():
|
||||
_ensure_not_greater(
|
||||
smaller=value,
|
||||
larger=total,
|
||||
smaller_path=value_path,
|
||||
larger_path=total_path,
|
||||
)
|
||||
for index, delay in enumerate(delays):
|
||||
_ensure_not_greater(
|
||||
smaller=delay,
|
||||
larger=total,
|
||||
smaller_path=f"{delays_path}[{index}]",
|
||||
larger_path=total_path,
|
||||
)
|
||||
|
||||
|
||||
def _request_policy(raw: Any, path: str) -> RequestRoutingPolicy:
|
||||
source = _mapping(raw, path)
|
||||
retry_delays = _field_delays(source, "retry_delays", path)
|
||||
request_timeout = _field_number(source, "request_timeout", path, minimum=1)
|
||||
total_timeout = _field_number(source, "total_timeout", path, minimum=1)
|
||||
retry_after_cap = _field_number(source, "retry_after_cap", path, maximum=3600)
|
||||
_validate_budget(
|
||||
total=total_timeout,
|
||||
total_path=f"{path}.total_timeout",
|
||||
bounded_values={
|
||||
f"{path}.request_timeout": request_timeout,
|
||||
f"{path}.retry_after_cap": retry_after_cap,
|
||||
},
|
||||
delays=retry_delays,
|
||||
delays_path=f"{path}.retry_delays",
|
||||
)
|
||||
return RequestRoutingPolicy(
|
||||
retry_delays=retry_delays,
|
||||
request_timeout=request_timeout,
|
||||
total_timeout=total_timeout,
|
||||
retry_after_cap=retry_after_cap,
|
||||
)
|
||||
|
||||
|
||||
def _text_policy(raw: Any) -> TextRoutingPolicy:
|
||||
path = "MODEL_ROUTING_POLICY.text"
|
||||
source = _mapping(raw, path)
|
||||
retry_delays = _field_delays(source, "retry_delays", path)
|
||||
request_timeout = _field_number(source, "request_timeout", path, minimum=1)
|
||||
stream_timeout = _field_number(source, "stream_timeout", path, minimum=1)
|
||||
total_timeout = _field_number(source, "total_timeout", path, minimum=1)
|
||||
retry_after_cap = _field_number(source, "retry_after_cap", path, maximum=3600)
|
||||
_validate_budget(
|
||||
total=total_timeout,
|
||||
total_path=f"{path}.total_timeout",
|
||||
bounded_values={
|
||||
f"{path}.request_timeout": request_timeout,
|
||||
f"{path}.stream_timeout": stream_timeout,
|
||||
f"{path}.retry_after_cap": retry_after_cap,
|
||||
},
|
||||
delays=retry_delays,
|
||||
delays_path=f"{path}.retry_delays",
|
||||
)
|
||||
return TextRoutingPolicy(
|
||||
retry_delays=retry_delays,
|
||||
request_timeout=request_timeout,
|
||||
stream_timeout=stream_timeout,
|
||||
total_timeout=total_timeout,
|
||||
retry_after_cap=retry_after_cap,
|
||||
)
|
||||
|
||||
|
||||
def _video_policy(raw: Any) -> VideoRoutingPolicy:
|
||||
path = "MODEL_ROUTING_POLICY.video"
|
||||
source = _mapping(raw, path)
|
||||
retry_delays = _field_delays(source, "submit_retry_delays", path)
|
||||
submit_timeout = _field_number(source, "submit_timeout", path, minimum=1)
|
||||
submit_total_timeout = _field_number(source, "submit_total_timeout", path, minimum=1)
|
||||
poll_request_timeout = _field_number(source, "poll_request_timeout", path, minimum=1)
|
||||
generation_timeout = _field_number(source, "generation_timeout", path, minimum=1)
|
||||
retry_after_cap = _field_number(source, "retry_after_cap", path, maximum=3600)
|
||||
_validate_budget(
|
||||
total=submit_total_timeout,
|
||||
total_path=f"{path}.submit_total_timeout",
|
||||
bounded_values={
|
||||
f"{path}.submit_timeout": submit_timeout,
|
||||
f"{path}.retry_after_cap": retry_after_cap,
|
||||
},
|
||||
delays=retry_delays,
|
||||
delays_path=f"{path}.submit_retry_delays",
|
||||
)
|
||||
_validate_budget(
|
||||
total=generation_timeout,
|
||||
total_path=f"{path}.generation_timeout",
|
||||
bounded_values={f"{path}.poll_request_timeout": poll_request_timeout},
|
||||
)
|
||||
return VideoRoutingPolicy(
|
||||
submit_retry_delays=retry_delays,
|
||||
submit_timeout=submit_timeout,
|
||||
submit_total_timeout=submit_total_timeout,
|
||||
poll_request_timeout=poll_request_timeout,
|
||||
generation_timeout=generation_timeout,
|
||||
retry_after_cap=retry_after_cap,
|
||||
)
|
||||
|
||||
|
||||
def load_model_routing_policy(raw_policy: Any | None = None) -> ModelRoutingPolicy:
|
||||
"""读取并校验模型路由策略,返回不可变对象。
|
||||
|
||||
``raw_policy`` 主要供测试和启动检查注入;省略时读取 Django settings。
|
||||
配置错误统一抛出带中文字段含义、错误值和合法范围的异常。
|
||||
"""
|
||||
|
||||
raw = settings.MODEL_ROUTING_POLICY if raw_policy is None else raw_policy
|
||||
source = _mapping(raw, "MODEL_ROUTING_POLICY")
|
||||
|
||||
path = "MODEL_ROUTING_POLICY"
|
||||
max_models = _integer(_required(source, "max_models", path), f"{path}.max_models", minimum=1, maximum=10)
|
||||
max_calls = _integer(_required(source, "max_calls", path), f"{path}.max_calls", minimum=1, maximum=20)
|
||||
if max_calls < max_models:
|
||||
raise _error(
|
||||
"MODEL_ROUTING_POLICY.max_calls",
|
||||
f"不得小于 max_models(最多尝试模型数量),当前为 {max_calls} < {max_models}",
|
||||
)
|
||||
jitter_ratio = _field_number(source, "jitter_ratio", path, maximum=1)
|
||||
postprocess_path = f"{path}.postprocess"
|
||||
postprocess_source = _mapping(_required(source, "postprocess", path), postprocess_path)
|
||||
|
||||
return ModelRoutingPolicy(
|
||||
max_models=max_models,
|
||||
max_calls=max_calls,
|
||||
jitter_ratio=jitter_ratio,
|
||||
text=_text_policy(_required(source, "text", path)),
|
||||
image=_request_policy(_required(source, "image", path), f"{path}.image"),
|
||||
audio=_request_policy(_required(source, "audio", path), f"{path}.audio"),
|
||||
video=_video_policy(_required(source, "video", path)),
|
||||
postprocess=PostprocessRoutingPolicy(
|
||||
retry_delays=_field_delays(postprocess_source, "retry_delays", postprocess_path)
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"ModelRoutingPolicy",
|
||||
"PostprocessRoutingPolicy",
|
||||
"RequestRoutingPolicy",
|
||||
"RoutingPolicyConfigurationError",
|
||||
"TextRoutingPolicy",
|
||||
"VideoRoutingPolicy",
|
||||
"load_model_routing_policy",
|
||||
]
|
||||
@@ -20,6 +20,7 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
import re
|
||||
from decimal import Decimal
|
||||
from functools import lru_cache
|
||||
from pathlib import Path
|
||||
|
||||
@@ -574,7 +575,7 @@ def stream_script_agent(
|
||||
):
|
||||
"""生成 SSE 帧字符串的同步生成器,供 StreamingHttpResponse 包裹。
|
||||
target_index 非空 = 精准只改第 N 镜(读全脚本上下文,后端强制保留其余镜原样)。"""
|
||||
from apps.ai.services import build_provider, create_ai_task
|
||||
from apps.ai.services import create_ai_task, stream_routed_text_request
|
||||
|
||||
yield _sse({"type": "tool", "id": "skill", "label": "加载电商脚本技能", "status": "running"})
|
||||
skill_loaded = bool(load_ecommerce_skill())
|
||||
@@ -619,6 +620,9 @@ def stream_script_agent(
|
||||
"mode": mode,
|
||||
"aspect_ratio": aspect_ratio,
|
||||
"total_duration": total_duration,
|
||||
"base_version_id": str(base_version_id or ""),
|
||||
"target_index": target_index,
|
||||
"model_routing_v1": True,
|
||||
},
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001 — 多为额度不足
|
||||
@@ -639,13 +643,45 @@ def stream_script_agent(
|
||||
task.status = AITask.Status.SUBMITTED
|
||||
task.submitted_at = timezone.now()
|
||||
task.save(update_fields=["status", "submitted_at", "updated_at"])
|
||||
provider = build_provider(model_config)
|
||||
for ev in provider.chat_completion_stream(
|
||||
model=model_config.name,
|
||||
endpoint=model_config.endpoint,
|
||||
def validate_script_text(raw_text: str) -> dict:
|
||||
candidate = normalize_draft(
|
||||
raw_text,
|
||||
aspect_ratio=aspect_ratio,
|
||||
total_duration=effective_duration,
|
||||
)
|
||||
if target_index is not None and base_draft:
|
||||
return _merge_single_segment(
|
||||
base_draft,
|
||||
candidate,
|
||||
target_index,
|
||||
aspect_ratio,
|
||||
effective_duration,
|
||||
)
|
||||
return candidate
|
||||
|
||||
routed_stream = stream_routed_text_request(
|
||||
task=task,
|
||||
primary_model=model_config,
|
||||
messages=messages,
|
||||
streaming=True,
|
||||
structured_output=True,
|
||||
business_operation="script_generate",
|
||||
temperature=0.85,
|
||||
):
|
||||
validate_text=validate_script_text,
|
||||
request_summary={
|
||||
"mode": mode,
|
||||
"target_index": target_index,
|
||||
"base_version_id": str(base_version_id or ""),
|
||||
"aspect_ratio": aspect_ratio,
|
||||
"total_duration": effective_duration,
|
||||
},
|
||||
)
|
||||
while True:
|
||||
try:
|
||||
ev = next(routed_stream)
|
||||
except StopIteration as completed:
|
||||
routed = completed.value
|
||||
break
|
||||
et = ev.get("type")
|
||||
if et == "reasoning":
|
||||
# 思考流:推理模型在出 JSON 前会先想很久,把思考逐字下发给前端(像对话一样可见),
|
||||
@@ -668,12 +704,8 @@ def stream_script_agent(
|
||||
if piece.strip():
|
||||
yield _sse({"type": "delta", "text": piece})
|
||||
elif et == "done":
|
||||
break
|
||||
raw = "".join(full)
|
||||
draft = normalize_draft(raw, aspect_ratio=aspect_ratio, total_duration=effective_duration)
|
||||
if target_index is not None and base_draft:
|
||||
# 精准改一镜:只采用新稿的第 target_index 镜,其余镜强制保持基准稿原样
|
||||
draft = _merge_single_segment(base_draft, draft, target_index, aspect_ratio, effective_duration)
|
||||
continue
|
||||
raw, _provider_response, draft = routed.value
|
||||
except Exception as exc: # noqa: BLE001
|
||||
_fail_task(task, reservation, str(exc))
|
||||
settled = True
|
||||
@@ -819,7 +851,7 @@ def regenerate_segment_via_agent(*, project, user, model_config: ModelConfig, se
|
||||
与 stream_script_agent 的 target_index 分支同源,但同步返回(不走 SSE)。计费 reserve→charge/release 闭环。"""
|
||||
from django.db import transaction
|
||||
|
||||
from apps.ai.services import build_provider, create_ai_task
|
||||
from apps.ai.services import create_ai_task, execute_routed_text_request
|
||||
from apps.billing.services.ledger import charge_reserved_credit
|
||||
|
||||
base_draft = _draft_from_version(segment.script_version) # 用 DB 行重建基准,别用可能 stale 的 content
|
||||
@@ -850,18 +882,40 @@ def regenerate_segment_via_agent(*, project, user, model_config: ModelConfig, se
|
||||
"endpoint": model_config.endpoint,
|
||||
"mode": "revise",
|
||||
"target_index": target_index,
|
||||
"model_routing_v1": True,
|
||||
},
|
||||
)
|
||||
reservation = task.credit_reservation
|
||||
# 每条真实尝试的平台成本由统一执行器累计;用户积分仍只结算这一条脚本任务。
|
||||
task.base_cost = Decimal("0")
|
||||
task.save(update_fields=["base_cost", "updated_at"])
|
||||
# 实际平台成本由每条 AIModelAttempt 累加;用户积分仍只结算这一条逻辑任务。
|
||||
task.base_cost = Decimal("0")
|
||||
task.save(update_fields=["base_cost", "updated_at"])
|
||||
try:
|
||||
task.status = AITask.Status.SUBMITTED
|
||||
task.submitted_at = timezone.now()
|
||||
task.save(update_fields=["status", "submitted_at", "updated_at"])
|
||||
provider = build_provider(model_config)
|
||||
response = provider.chat_completion(model=model_config.name, endpoint=model_config.endpoint, messages=messages)
|
||||
raw = provider.extract_text(response)
|
||||
draft = normalize_draft(raw, aspect_ratio=aspect_ratio, total_duration=total_duration)
|
||||
draft = _merge_single_segment(base_draft, draft, target_index, aspect_ratio, total_duration)
|
||||
def validate_segment_text(raw_text: str) -> dict:
|
||||
candidate = normalize_draft(raw_text, aspect_ratio=aspect_ratio, total_duration=total_duration)
|
||||
return _merge_single_segment(base_draft, candidate, target_index, aspect_ratio, total_duration)
|
||||
|
||||
routed = execute_routed_text_request(
|
||||
task=task,
|
||||
primary_model=model_config,
|
||||
messages=messages,
|
||||
streaming=False,
|
||||
structured_output=True,
|
||||
business_operation="script_generate",
|
||||
temperature=0.3,
|
||||
validate_text=validate_segment_text,
|
||||
request_summary={
|
||||
"mode": "revise",
|
||||
"target_index": target_index,
|
||||
"base_version_id": str(segment.script_version_id),
|
||||
},
|
||||
)
|
||||
raw, _response, draft = routed.value
|
||||
with transaction.atomic():
|
||||
task.status = AITask.Status.SUCCEEDED
|
||||
task.response_payload = {"raw": raw[:8000]}
|
||||
@@ -871,6 +925,6 @@ def regenerate_segment_via_agent(*, project, user, model_config: ModelConfig, se
|
||||
charge_reserved_credit(reservation=reservation, actual_amount=task.actual_cost)
|
||||
script = persist_script_draft(project=project, user=user, task=task, draft=draft, source="revise")
|
||||
return script
|
||||
except Exception:
|
||||
_fail_task(task, reservation, "单镜重跑失败")
|
||||
except Exception as exc:
|
||||
_fail_task(task, reservation, str(exc) or "单镜重跑失败")
|
||||
raise
|
||||
|
||||
+891
-109
File diff suppressed because it is too large
Load Diff
@@ -63,7 +63,7 @@ def generate_model_triview_task(self, task_id: str) -> str:
|
||||
@app.task(bind=True, max_retries=0)
|
||||
def poll_free_video_task(self, task_id: str, attempt: int = 0) -> str:
|
||||
"""自由创作视频·worker 兜底轮询:每 30s 一次自重排(不依赖 celery beat),
|
||||
上限 60 次(≈30 分钟,足够 Seedance 5-10 分钟出片)。finalize 幂等(POSTPROCESSING 认领),
|
||||
总等待上限由统一视频 generation_timeout 配置控制。finalize 幂等(POSTPROCESSING 认领),
|
||||
与前端主动 poll 并存不双扣。轮询本身出错不重试(max_retries=0),下一次自重排继续。"""
|
||||
from apps.ai.free_video import finalize_free_video
|
||||
from apps.ai.models import AITask
|
||||
@@ -82,7 +82,9 @@ def poll_free_video_task(self, task_id: str, attempt: int = 0) -> str:
|
||||
|
||||
if getattr(dj_settings, "CELERY_TASK_ALWAYS_EAGER", False):
|
||||
return task_id
|
||||
if task.status in (AITask.Status.SUBMITTED, AITask.Status.POLLING) and attempt < 60:
|
||||
# 是否终止由 finalize_free_video 读取统一 generation_timeout 决定;这里不再维护一套
|
||||
# “60 次轮询”硬编码,避免修改成片等待时限时 Worker 与 Web 轮询口径不一致。
|
||||
if task.status in (AITask.Status.SUBMITTED, AITask.Status.POLLING):
|
||||
poll_free_video_task.apply_async(args=[task_id, attempt + 1], countdown=30)
|
||||
return task_id
|
||||
|
||||
|
||||
@@ -0,0 +1,248 @@
|
||||
from decimal import Decimal
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import requests
|
||||
from django.test import TestCase
|
||||
|
||||
from apps.accounts.models import Team, TeamMember, User
|
||||
from apps.ai.models import AITask, ModelConfig, ModelProvider
|
||||
from apps.ai.services import generate_base_asset, run_base_asset_task
|
||||
from apps.assets.models import Asset, AssetFile
|
||||
from apps.billing.models import CreditAccount, CreditLedger
|
||||
from apps.products.models import Product
|
||||
from apps.projects.models import BaseAssetGroup, Project
|
||||
|
||||
|
||||
def _image_metadata(*, outbound=True, priority=20, base_cost="0.50"):
|
||||
return {
|
||||
"routing": {
|
||||
"fallback_on_failure": outbound,
|
||||
"fallback_candidate": True,
|
||||
},
|
||||
"capabilities": {
|
||||
"operations": ["image_generate", "image_edit"],
|
||||
"features": [],
|
||||
"reference_modes": ["none", "single", "multiple"],
|
||||
"max_reference_images": 9,
|
||||
"aspect_ratios": ["1:1", "3:4", "4:5", "9:16", "16:9"],
|
||||
},
|
||||
"pricing": {"base_cost_yuan": base_cost},
|
||||
"test_priority": priority,
|
||||
}
|
||||
|
||||
|
||||
class BaseAssetRoutingTests(TestCase):
|
||||
def setUp(self):
|
||||
ModelConfig.objects.filter(capability=ModelConfig.Capability.IMAGE).update(
|
||||
status=ModelConfig.Status.DISABLED
|
||||
)
|
||||
self.user = User.objects.create_user(username="base-routing", password="x")
|
||||
self.team = Team.objects.create(name="Base Routing", owner=self.user)
|
||||
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
|
||||
CreditAccount.objects.create(team=self.team, balance=Decimal("1000"))
|
||||
self.product = Product.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
title="橙色童装上衣",
|
||||
category="服饰内衣",
|
||||
)
|
||||
self.project = Project.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
product=self.product,
|
||||
name="基础资产路由项目",
|
||||
)
|
||||
self.provider_mocks = {}
|
||||
patch("apps.ai.services.get_image_provider", side_effect=self._provider_for).start()
|
||||
patch("apps.ai.services._store_generated_media", side_effect=self._store_media).start()
|
||||
patch("apps.ai.tasks.generate_base_asset_task.delay").start()
|
||||
self.addCleanup(patch.stopall)
|
||||
|
||||
def provider(self, name, priority):
|
||||
return ModelProvider.objects.create(
|
||||
name=name,
|
||||
display_name=name,
|
||||
status=ModelProvider.Status.ACTIVE,
|
||||
metadata={"routing": {"fallback_priority": priority}},
|
||||
)
|
||||
|
||||
def model(self, provider, name, *, outbound=True, base_cost="0.50"):
|
||||
return ModelConfig.objects.create(
|
||||
provider=provider,
|
||||
name=name,
|
||||
display_name=name,
|
||||
capability=ModelConfig.Capability.IMAGE,
|
||||
endpoint="images/generations",
|
||||
unit_price=Decimal("20"),
|
||||
status=ModelConfig.Status.ACTIVE,
|
||||
metadata=_image_metadata(
|
||||
outbound=outbound,
|
||||
priority=(provider.metadata or {}).get("routing", {}).get("fallback_priority", 20),
|
||||
base_cost=base_cost,
|
||||
),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _new_provider_mock():
|
||||
provider = Mock()
|
||||
provider.image_generation.return_value = {"data": [{"url": "http://example.test/out.png"}]}
|
||||
provider.image_edit.return_value = {"data": [{"url": "http://example.test/out.png"}]}
|
||||
provider.extract_first_media_url.side_effect = lambda response: response["data"][0]["url"]
|
||||
return provider
|
||||
|
||||
@staticmethod
|
||||
def _new_generation_only_provider_mock():
|
||||
provider = Mock(spec=["image_generation", "extract_first_media_url"])
|
||||
provider.image_generation.return_value = {"data": [{"url": "http://example.test/out.png"}]}
|
||||
provider.extract_first_media_url.side_effect = lambda response: response["data"][0]["url"]
|
||||
return provider
|
||||
|
||||
def _provider_for(self, model):
|
||||
return self.provider_mocks.setdefault(model.id, self._new_provider_mock())
|
||||
|
||||
def _store_media(self, **kwargs):
|
||||
return Asset.objects.create(
|
||||
team=kwargs["team"],
|
||||
created_by=kwargs["user"],
|
||||
name=kwargs["name"],
|
||||
asset_type=kwargs["asset_type"],
|
||||
source=Asset.Source.AI_GENERATED,
|
||||
category=kwargs["category"],
|
||||
origin_task=kwargs["task"],
|
||||
)
|
||||
|
||||
def reference_asset(self, name="商品主图", category=Asset.Category.PRODUCT_IMAGE):
|
||||
asset = Asset.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
name=name,
|
||||
asset_type=Asset.Type.IMAGE,
|
||||
source=Asset.Source.UPLOAD,
|
||||
category=category,
|
||||
)
|
||||
AssetFile.objects.create(
|
||||
asset=asset,
|
||||
object_key=f"{name}.png",
|
||||
bucket="bucket",
|
||||
content_type="image/png",
|
||||
preview_url=f"http://example.test/{name}.png",
|
||||
is_primary=True,
|
||||
)
|
||||
return asset
|
||||
|
||||
def submit(self, kind, *, reference_asset_id=None):
|
||||
return generate_base_asset(
|
||||
project=self.project,
|
||||
user=self.user,
|
||||
kind=kind,
|
||||
prompt=f"{kind} 基础资产提示词",
|
||||
label=f"{kind}-label",
|
||||
reference_asset_id=reference_asset_id,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def ledger_count(task, ledger_type):
|
||||
return CreditLedger.objects.filter(task=task, ledger_type=ledger_type).count()
|
||||
|
||||
def test_product_reference_uses_routed_image_edit_and_keeps_group(self):
|
||||
primary = self.model(self.provider("base-product-primary", 20), "base-product")
|
||||
cover = self.reference_asset()
|
||||
self.product.cover_asset = cover
|
||||
self.product.save(update_fields=["cover_asset"])
|
||||
task = self.submit(BaseAssetGroup.Kind.PRODUCT)
|
||||
|
||||
run_base_asset_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertTrue(task.request_payload["model_routing_v1"])
|
||||
attempt = task.model_attempts.get()
|
||||
self.assertEqual(attempt.model_config_id, primary.id)
|
||||
self.assertEqual(attempt.operation, "image_edit")
|
||||
self.assertEqual(attempt.request_summary["reference_images"], 1)
|
||||
self.assertEqual(attempt.request_summary["base_asset_kind"], "product")
|
||||
call = self.provider_mocks[primary.id].image_edit.call_args
|
||||
self.assertEqual(call.kwargs["images"], ["http://example.test/商品主图.png"])
|
||||
group = BaseAssetGroup.objects.get(project=self.project, kind=BaseAssetGroup.Kind.PRODUCT)
|
||||
self.assertEqual(group.adopted_asset.origin_task_id, task.id)
|
||||
self.assertEqual(group.metadata["label"], "product-label")
|
||||
|
||||
def test_person_and_scene_without_reference_use_routed_generation(self):
|
||||
primary = self.model(self.provider("base-generate-primary", 20), "base-generate")
|
||||
for kind in (BaseAssetGroup.Kind.PERSON, BaseAssetGroup.Kind.SCENE):
|
||||
with self.subTest(kind=kind):
|
||||
task = self.submit(kind)
|
||||
run_base_asset_task(task_id=str(task.id))
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
attempt = task.model_attempts.get()
|
||||
self.assertEqual(attempt.model_config_id, primary.id)
|
||||
self.assertEqual(attempt.operation, "image_generate")
|
||||
self.assertEqual(attempt.request_summary["reference_images"], 0)
|
||||
self.assertEqual(attempt.request_summary["base_asset_kind"], kind)
|
||||
group = BaseAssetGroup.objects.get(project=self.project, kind=kind)
|
||||
self.assertEqual(group.adopted_asset.origin_task_id, task.id)
|
||||
|
||||
def test_person_rerun_reference_is_preserved_across_fallback(self):
|
||||
primary = self.model(
|
||||
self.provider("base-person-primary", 100), "person-primary", base_cost="0.25"
|
||||
)
|
||||
candidate = self.model(
|
||||
self.provider("volcano", 10), "person-candidate", outbound=False, base_cost="0.75"
|
||||
)
|
||||
portrait = self.reference_asset("人物立绘", Asset.Category.PERSON)
|
||||
task = self.submit(BaseAssetGroup.Kind.PERSON, reference_asset_id=str(portrait.id))
|
||||
self.provider_mocks[primary.id] = self._new_provider_mock()
|
||||
self.provider_mocks[primary.id].image_edit.side_effect = requests.ConnectionError("offline")
|
||||
|
||||
run_base_asset_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
attempts = list(task.model_attempts.all())
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertEqual([a.model_config_id for a in attempts], [primary.id, primary.id, candidate.id])
|
||||
self.assertTrue(attempts[-1].is_fallback)
|
||||
self.assertTrue(all(a.request_summary["reference_images"] == 1 for a in attempts))
|
||||
self.assertTrue(all(a.public_model_name == "AirShelf Image" for a in attempts))
|
||||
candidate_call = self.provider_mocks[candidate.id].image_edit.call_args
|
||||
self.assertEqual(candidate_call.kwargs["images"], ["http://example.test/人物立绘.png"])
|
||||
self.assertEqual(task.base_cost, Decimal("0.7500"))
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
|
||||
|
||||
def test_all_base_asset_candidates_fail_release_once(self):
|
||||
primary = self.model(self.provider("base-all-primary", 100), "base-primary")
|
||||
candidate = self.model(self.provider("volcano", 10), "base-candidate", outbound=False)
|
||||
task = self.submit(BaseAssetGroup.Kind.SCENE)
|
||||
for model in (primary, candidate):
|
||||
self.provider_mocks[model.id] = self._new_provider_mock()
|
||||
self.provider_mocks[model.id].image_generation.side_effect = requests.ConnectionError("offline")
|
||||
|
||||
run_base_asset_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.FAILED)
|
||||
self.assertEqual(task.model_attempts.count(), 3)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
|
||||
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
|
||||
|
||||
def test_direct_product_primary_uses_image_input_and_public_real_name(self):
|
||||
primary = self.model(
|
||||
self.provider("volcano", 10), "seedream-base-product", outbound=False
|
||||
)
|
||||
self.provider_mocks[primary.id] = self._new_generation_only_provider_mock()
|
||||
cover = self.reference_asset()
|
||||
self.product.cover_asset = cover
|
||||
self.product.save(update_fields=["cover_asset"])
|
||||
task = self.submit(BaseAssetGroup.Kind.PRODUCT)
|
||||
|
||||
run_base_asset_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
call = self.provider_mocks[primary.id].image_generation.call_args
|
||||
self.assertEqual(call.kwargs["image"], ["http://example.test/商品主图.png"])
|
||||
self.assertEqual(task.model_attempts.get().public_model_name, primary.display_name)
|
||||
@@ -0,0 +1,238 @@
|
||||
from decimal import Decimal
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import requests
|
||||
from django.test import TestCase
|
||||
|
||||
from apps.accounts.models import Team, TeamMember, User
|
||||
from apps.ai.models import AITask, ModelConfig, ModelProvider
|
||||
from apps.ai.services import EXTRACT_TEXT_MODEL_NAME, run_extract_entities_task, submit_extract_entities
|
||||
from apps.billing.models import CreditAccount, CreditLedger
|
||||
from apps.products.models import Product
|
||||
from apps.projects.models import Project, ScriptSegment, ScriptVersion
|
||||
|
||||
|
||||
def _metadata(*, outbound=True, base_cost="0.50"):
|
||||
return {
|
||||
"routing": {"fallback_on_failure": outbound, "fallback_candidate": True},
|
||||
"capabilities": {
|
||||
"operations": ["chat"],
|
||||
"features": ["streaming", "structured_output"],
|
||||
},
|
||||
"pricing": {"base_cost_yuan": base_cost},
|
||||
}
|
||||
|
||||
|
||||
class EntityExtractionRoutingTests(TestCase):
|
||||
def setUp(self):
|
||||
ModelConfig.objects.filter(capability=ModelConfig.Capability.TEXT).update(
|
||||
status=ModelConfig.Status.DISABLED
|
||||
)
|
||||
self.user = User.objects.create_user(username="entity-routing", password="x")
|
||||
self.team = Team.objects.create(name="Entity Routing", owner=self.user)
|
||||
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
|
||||
CreditAccount.objects.create(team=self.team, balance=Decimal("1000"))
|
||||
product = Product.objects.create(team=self.team, created_by=self.user, title="测试商品")
|
||||
self.project = Project.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
product=product,
|
||||
name="实体提取路由项目",
|
||||
metadata={"cast": ["旧角色"], "entities_extracted": False},
|
||||
)
|
||||
self.script = ScriptVersion.objects.create(
|
||||
project=self.project,
|
||||
title="脚本",
|
||||
content="结构化脚本",
|
||||
is_adopted=True,
|
||||
)
|
||||
self.segment = ScriptSegment.objects.create(
|
||||
script_version=self.script,
|
||||
sort_order=0,
|
||||
narration="女主在客厅展示商品",
|
||||
visual_prompt="女主站在客厅",
|
||||
entity_refs=["old"],
|
||||
)
|
||||
self.provider_mocks = {}
|
||||
patch("apps.ai.services.get_text_provider", side_effect=self._provider_for).start()
|
||||
patch("apps.ai.tasks.extract_entities_task.delay").start()
|
||||
self.addCleanup(patch.stopall)
|
||||
|
||||
def provider(self, name, priority):
|
||||
return ModelProvider.objects.create(
|
||||
name=name,
|
||||
display_name=name,
|
||||
status=ModelProvider.Status.ACTIVE,
|
||||
metadata={"routing": {"fallback_priority": priority}},
|
||||
)
|
||||
|
||||
def model(self, provider, name, *, outbound=True, base_cost="0.50"):
|
||||
return ModelConfig.objects.create(
|
||||
provider=provider,
|
||||
name=name,
|
||||
display_name=name,
|
||||
capability=ModelConfig.Capability.TEXT,
|
||||
endpoint="chat/completions",
|
||||
unit_price=Decimal("10"),
|
||||
status=ModelConfig.Status.ACTIVE,
|
||||
metadata=_metadata(outbound=outbound, base_cost=base_cost),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def valid_events():
|
||||
return [
|
||||
{
|
||||
"type": "delta",
|
||||
"text": (
|
||||
'{"entities":['
|
||||
'{"id":"c1","type":"character","name":"女主","visual_prompt":"都市女主"},'
|
||||
'{"id":"s1","type":"scene","name":"客厅","visual_prompt":"现代客厅"}'
|
||||
'],"segments":[{"index":0,"entity_refs":["c1","s1"]}]}'
|
||||
),
|
||||
},
|
||||
{"type": "done"},
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def invalid_events():
|
||||
return [{"type": "delta", "text": "这是一段没有 JSON 的解释"}, {"type": "done"}]
|
||||
|
||||
@classmethod
|
||||
def _new_provider_mock(cls):
|
||||
provider = Mock()
|
||||
provider.chat_completion_stream.return_value = cls.valid_events()
|
||||
return provider
|
||||
|
||||
def _provider_for(self, model):
|
||||
return self.provider_mocks.setdefault(model.id, self._new_provider_mock())
|
||||
|
||||
def submit(self):
|
||||
return submit_extract_entities(project=self.project, user=self.user)
|
||||
|
||||
@staticmethod
|
||||
def ledger_count(task, ledger_type):
|
||||
return CreditLedger.objects.filter(task=task, ledger_type=ledger_type).count()
|
||||
|
||||
def test_first_success_records_streaming_structured_attempt_and_persists_entities(self):
|
||||
primary = self.model(self.provider("entity-primary", 20), "entity-primary")
|
||||
task = self.submit()
|
||||
|
||||
run_extract_entities_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
self.project.refresh_from_db()
|
||||
self.segment.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertTrue(task.request_payload["model_routing_v1"])
|
||||
attempt = task.model_attempts.get()
|
||||
self.assertEqual(attempt.model_config_id, primary.id)
|
||||
self.assertEqual(attempt.operation, "chat")
|
||||
self.assertTrue(attempt.request_summary["streaming"])
|
||||
self.assertTrue(attempt.request_summary["structured_output"])
|
||||
self.assertEqual(attempt.request_summary["business_operation"], "entity_extract")
|
||||
self.assertEqual(attempt.public_model_name, "AirShelf Script")
|
||||
self.assertTrue(self.project.metadata["entities_extracted"])
|
||||
self.assertEqual(self.project.metadata["cast"], ["女主"])
|
||||
self.assertEqual(self.project.metadata["scenes"], ["客厅"])
|
||||
self.assertEqual(self.segment.entity_refs, ["c1", "s1"])
|
||||
call = self.provider_mocks[primary.id].chat_completion_stream.call_args
|
||||
self.assertGreater(call.kwargs["timeout"], 0)
|
||||
self.assertEqual(call.kwargs["temperature"], 0.3)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
|
||||
def test_invalid_structure_retries_then_falls_back_and_records_failed_call_costs(self):
|
||||
primary = self.model(
|
||||
self.provider("entity-fallback-primary", 100),
|
||||
"entity-primary",
|
||||
base_cost="0.25",
|
||||
)
|
||||
candidate = self.model(
|
||||
self.provider("entity-fallback-candidate", 10),
|
||||
"entity-candidate",
|
||||
outbound=False,
|
||||
base_cost="0.75",
|
||||
)
|
||||
primary_mock = self._new_provider_mock()
|
||||
primary_mock.chat_completion_stream.side_effect = [
|
||||
self.invalid_events(),
|
||||
self.invalid_events(),
|
||||
self.invalid_events(),
|
||||
]
|
||||
self.provider_mocks[primary.id] = primary_mock
|
||||
task = self.submit()
|
||||
|
||||
run_extract_entities_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
attempts = list(task.model_attempts.all())
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertEqual(
|
||||
[a.model_config_id for a in attempts],
|
||||
[primary.id, primary.id, primary.id, candidate.id],
|
||||
)
|
||||
self.assertEqual(
|
||||
[a.status for a in attempts],
|
||||
["failed", "failed", "failed", "succeeded"],
|
||||
)
|
||||
self.assertTrue(attempts[-1].is_fallback)
|
||||
self.assertTrue(attempts[0].response_summary["validation_failed"])
|
||||
self.assertEqual(task.base_cost, Decimal("1.5000"))
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
|
||||
|
||||
def test_all_candidates_fail_releases_and_preserves_old_project_metadata(self):
|
||||
primary = self.model(self.provider("entity-all-primary", 100), "entity-primary")
|
||||
fallback_1 = self.model(
|
||||
self.provider("entity-all-fallback-1", 10), "entity-fallback-1", outbound=False
|
||||
)
|
||||
fallback_2 = self.model(
|
||||
self.provider("entity-all-fallback-2", 20), "entity-fallback-2", outbound=False
|
||||
)
|
||||
for model in (primary, fallback_1, fallback_2):
|
||||
provider = self._new_provider_mock()
|
||||
provider.chat_completion_stream.side_effect = requests.ConnectionError("offline")
|
||||
self.provider_mocks[model.id] = provider
|
||||
task = self.submit()
|
||||
|
||||
run_extract_entities_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
self.project.refresh_from_db()
|
||||
self.segment.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.FAILED)
|
||||
self.assertEqual(task.model_attempts.count(), 4)
|
||||
self.assertEqual(self.project.metadata["cast"], ["旧角色"])
|
||||
self.assertFalse(self.project.metadata["entities_extracted"])
|
||||
self.assertEqual(self.segment.entity_refs, ["old"])
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
|
||||
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
|
||||
|
||||
def test_direct_doubao_primary_retries_but_does_not_switch_when_outbound_disabled(self):
|
||||
primary = self.model(
|
||||
self.provider("doubao", 10), EXTRACT_TEXT_MODEL_NAME, outbound=False
|
||||
)
|
||||
candidate = self.model(
|
||||
self.provider("entity-unused-candidate", 20), "entity-unused", outbound=False
|
||||
)
|
||||
provider = self._new_provider_mock()
|
||||
provider.chat_completion_stream.side_effect = requests.ConnectionError("offline")
|
||||
self.provider_mocks[primary.id] = provider
|
||||
task = self.submit()
|
||||
|
||||
run_extract_entities_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
attempts = list(task.model_attempts.all())
|
||||
self.assertEqual(task.status, AITask.Status.FAILED)
|
||||
self.assertEqual(
|
||||
[attempt.model_config_id for attempt in attempts],
|
||||
[primary.id, primary.id, primary.id],
|
||||
)
|
||||
self.assertFalse(any(attempt.model_config_id == candidate.id for attempt in attempts))
|
||||
self.assertTrue(all(attempt.public_model_name == primary.display_name for attempt in attempts))
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
|
||||
@@ -0,0 +1,243 @@
|
||||
from copy import deepcopy
|
||||
from datetime import timedelta
|
||||
from decimal import Decimal
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import requests
|
||||
from django.conf import settings
|
||||
from django.test import TestCase, override_settings
|
||||
from django.utils import timezone
|
||||
|
||||
from apps.accounts.models import Team, TeamMember, User
|
||||
from apps.ai.free_video import finalize_free_video, submit_free_video
|
||||
from apps.ai.models import AITask, ModelConfig, ModelProvider
|
||||
from apps.assets.models import Asset
|
||||
from apps.billing.models import CreditAccount, CreditLedger
|
||||
|
||||
|
||||
STANDARD = "doubao-seedance-2-0-260128"
|
||||
|
||||
|
||||
def _fast_routing_policy():
|
||||
policy = deepcopy(settings.MODEL_ROUTING_POLICY)
|
||||
policy["jitter_ratio"] = 0
|
||||
policy["video"]["submit_retry_delays"] = [0]
|
||||
return policy
|
||||
|
||||
|
||||
def _candidate_metadata(*, outbound=False, max_images=9, max_videos=3, max_audios=3):
|
||||
return {
|
||||
"routing": {
|
||||
"fallback_on_failure": outbound,
|
||||
"fallback_candidate": True,
|
||||
},
|
||||
"capabilities": {
|
||||
"operations": ["video_generate"],
|
||||
"features": ["generate_audio"],
|
||||
"max_reference_images": max_images,
|
||||
"max_reference_videos": max_videos,
|
||||
"max_reference_audios": max_audios,
|
||||
"aspect_ratios": ["16:9"],
|
||||
"resolutions": ["480p"],
|
||||
"durations": [4],
|
||||
},
|
||||
"pricing": {
|
||||
"unit": "cny_per_million_tokens",
|
||||
"default": {"no_ref_video": 23, "with_ref_video": 14},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@override_settings(MODEL_ROUTING_POLICY=_fast_routing_policy())
|
||||
class FreeVideoRoutingTests(TestCase):
|
||||
def setUp(self):
|
||||
self.user = User.objects.create_user(username="free-video-routing", password="x")
|
||||
self.team = Team.objects.create(name="Free Video Routing", owner=self.user)
|
||||
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
|
||||
CreditAccount.objects.create(team=self.team, balance=Decimal("10000"))
|
||||
ModelConfig.objects.filter(capability=ModelConfig.Capability.VIDEO).update(
|
||||
status=ModelConfig.Status.DISABLED,
|
||||
is_default=False,
|
||||
)
|
||||
self.primary = ModelConfig.objects.select_related("provider").get(
|
||||
name=STANDARD,
|
||||
capability=ModelConfig.Capability.VIDEO,
|
||||
)
|
||||
primary_metadata = deepcopy(self.primary.metadata or {})
|
||||
primary_metadata.setdefault("routing", {})["fallback_on_failure"] = True
|
||||
primary_metadata["routing"]["fallback_candidate"] = True
|
||||
self.primary.metadata = primary_metadata
|
||||
self.primary.status = ModelConfig.Status.ACTIVE
|
||||
self.primary.is_default = True
|
||||
self.primary.provider.status = ModelProvider.Status.ACTIVE
|
||||
self.primary.provider.metadata = {"routing": {"fallback_priority": 10}}
|
||||
self.primary.provider.save(update_fields=["status", "metadata", "updated_at"])
|
||||
self.primary.save(update_fields=["metadata", "status", "is_default", "updated_at"])
|
||||
self.provider_mocks = {}
|
||||
patch("apps.ai.services.get_video_provider", side_effect=self._provider_for).start()
|
||||
patch("apps.ai.tasks.poll_free_video_task.apply_async").start()
|
||||
patch("apps.ai.free_video._notify_failure").start()
|
||||
self.addCleanup(patch.stopall)
|
||||
|
||||
def candidate(self, name="free-video-candidate", **metadata_overrides):
|
||||
provider = ModelProvider.objects.create(
|
||||
name=f"{name}-provider",
|
||||
display_name=name,
|
||||
status=ModelProvider.Status.ACTIVE,
|
||||
base_url="https://video.example/v1",
|
||||
metadata={"routing": {"fallback_priority": 20}},
|
||||
)
|
||||
metadata = _candidate_metadata(**metadata_overrides)
|
||||
return ModelConfig.objects.create(
|
||||
provider=provider,
|
||||
name=name,
|
||||
display_name=name,
|
||||
capability=ModelConfig.Capability.VIDEO,
|
||||
endpoint="videos",
|
||||
status=ModelConfig.Status.ACTIVE,
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
def _provider_for(self, model):
|
||||
return self.provider_mocks.setdefault(model.id, Mock())
|
||||
|
||||
def params(self, **overrides):
|
||||
params = {
|
||||
"prompt": "A product video",
|
||||
"mode": "universal",
|
||||
"model": STANDARD,
|
||||
"aspect_ratio": "16:9",
|
||||
"resolution": "480p",
|
||||
"duration": 4,
|
||||
"generate_audio": True,
|
||||
"references": [],
|
||||
}
|
||||
params.update(overrides)
|
||||
return params
|
||||
|
||||
@staticmethod
|
||||
def ledger_count(task, ledger_type):
|
||||
return CreditLedger.objects.filter(task=task, ledger_type=ledger_type).count()
|
||||
|
||||
def test_multi_material_fallback_uses_actual_model_for_poll_and_settlement(self):
|
||||
candidate = self.candidate()
|
||||
primary_provider = self._provider_for(self.primary)
|
||||
fallback_provider = self._provider_for(candidate)
|
||||
primary_provider.create_video_task.side_effect = requests.ConnectionError("primary offline")
|
||||
fallback_provider.create_video_task.return_value = {"id": "free-fallback", "status": "queued"}
|
||||
references = [
|
||||
{"url": "https://media.example/ref.png", "type": "image", "label": "image"},
|
||||
{"url": "https://media.example/ref.mp4", "type": "video", "label": "video", "duration": 1},
|
||||
{"url": "https://media.example/ref.mp3", "type": "audio", "label": "audio"},
|
||||
]
|
||||
|
||||
task = submit_free_video(team=self.team, user=self.user, params=self.params(references=references))
|
||||
|
||||
self.assertEqual(task.status, AITask.Status.SUBMITTED)
|
||||
self.assertEqual(task.provider_task_id, "free-fallback")
|
||||
self.assertEqual(task.request_payload["actual_model_config_id"], str(candidate.id))
|
||||
attempts = list(task.model_attempts.order_by("sequence"))
|
||||
self.assertEqual(
|
||||
[attempt.model_config_id for attempt in attempts],
|
||||
[self.primary.id, self.primary.id, candidate.id],
|
||||
)
|
||||
self.assertEqual(attempts[-1].request_summary["reference_images"], 1)
|
||||
self.assertEqual(attempts[-1].request_summary["reference_videos"], 1)
|
||||
self.assertEqual(attempts[-1].request_summary["reference_audios"], 1)
|
||||
submit_call = fallback_provider.create_video_task.call_args.kwargs
|
||||
self.assertEqual(len(submit_call["content_items"]), 3)
|
||||
self.assertEqual(submit_call["timeout"], 120.0)
|
||||
|
||||
fallback_provider.poll_video_task.return_value = {
|
||||
"status": "succeeded",
|
||||
"usage": {"total_tokens": 30000},
|
||||
"content": {"video_url": "https://video.example/result.mp4"},
|
||||
}
|
||||
fallback_provider.extract_first_media_url.return_value = "https://video.example/result.mp4"
|
||||
with patch("apps.ai.free_video._store_free_video_media") as store_media:
|
||||
store_media.return_value = Asset.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
name="free result",
|
||||
asset_type=Asset.Type.VIDEO,
|
||||
source=Asset.Source.AI_GENERATED,
|
||||
category=Asset.Category.FREE_CREATE,
|
||||
origin_task=task,
|
||||
)
|
||||
result = finalize_free_video(task=task)
|
||||
|
||||
self.assertEqual(result.status, AITask.Status.SUCCEEDED)
|
||||
self.assertFalse(primary_provider.poll_video_task.called)
|
||||
fallback_provider.poll_video_task.assert_called_once_with(
|
||||
endpoint=candidate.endpoint,
|
||||
provider_task_id="free-fallback",
|
||||
timeout=60.0,
|
||||
)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
|
||||
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
|
||||
|
||||
def test_candidate_missing_audio_capacity_is_filtered_out(self):
|
||||
candidate = self.candidate("free-video-incompatible", max_audios=0)
|
||||
primary_provider = self._provider_for(self.primary)
|
||||
unused_provider = self._provider_for(candidate)
|
||||
primary_provider.create_video_task.side_effect = requests.ConnectionError("primary offline")
|
||||
references = [
|
||||
{"url": "https://media.example/ref.png", "type": "image"},
|
||||
{"url": "https://media.example/ref.mp3", "type": "audio"},
|
||||
]
|
||||
|
||||
task = submit_free_video(team=self.team, user=self.user, params=self.params(references=references))
|
||||
|
||||
self.assertEqual(task.status, AITask.Status.FAILED)
|
||||
self.assertFalse(unused_provider.create_video_task.called)
|
||||
self.assertEqual(task.model_attempts.count(), 2)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
|
||||
|
||||
def test_submit_read_timeout_is_not_retried_or_fallbacked(self):
|
||||
candidate = self.candidate("free-video-timeout-unused")
|
||||
primary_provider = self._provider_for(self.primary)
|
||||
unused_provider = self._provider_for(candidate)
|
||||
primary_provider.create_video_task.side_effect = requests.ReadTimeout("response lost")
|
||||
|
||||
task = submit_free_video(team=self.team, user=self.user, params=self.params())
|
||||
|
||||
self.assertEqual(task.status, AITask.Status.FAILED)
|
||||
self.assertEqual(task.model_attempts.count(), 1)
|
||||
self.assertFalse(unused_provider.create_video_task.called)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
|
||||
|
||||
def test_poll_network_error_keeps_inflight_task_and_reservation(self):
|
||||
provider = self._provider_for(self.primary)
|
||||
provider.create_video_task.return_value = {"id": "free-poll", "status": "queued"}
|
||||
task = submit_free_video(team=self.team, user=self.user, params=self.params())
|
||||
provider.poll_video_task.side_effect = requests.ConnectionError("poll offline")
|
||||
|
||||
with self.assertRaises(requests.ConnectionError):
|
||||
finalize_free_video(task=task)
|
||||
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.SUBMITTED)
|
||||
self.assertEqual(task.model_attempts.count(), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
|
||||
self.assertGreater(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
|
||||
|
||||
def test_generation_timeout_releases_without_polling_or_resubmitting(self):
|
||||
provider = self._provider_for(self.primary)
|
||||
provider.create_video_task.return_value = {"id": "free-slow", "status": "queued"}
|
||||
task = submit_free_video(team=self.team, user=self.user, params=self.params())
|
||||
task.submitted_at = timezone.now() - timedelta(seconds=1900)
|
||||
task.save(update_fields=["submitted_at", "updated_at"])
|
||||
|
||||
result = finalize_free_video(task=task)
|
||||
|
||||
result.refresh_from_db()
|
||||
self.assertEqual(result.status, AITask.Status.FAILED)
|
||||
self.assertFalse(provider.poll_video_task.called)
|
||||
self.assertEqual(provider.create_video_task.call_count, 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
|
||||
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
|
||||
@@ -0,0 +1,348 @@
|
||||
from copy import deepcopy
|
||||
from datetime import timedelta
|
||||
|
||||
from django.conf import settings
|
||||
from django.test import SimpleTestCase, TestCase, override_settings
|
||||
from django.utils import timezone
|
||||
|
||||
from apps.adminpanel.serializers import AdminModelConfigSerializer, AdminModelProviderSerializer
|
||||
from apps.ai.catalog import VOLCANO_MODELS, VOLCANO_PROVIDER, YUNQI_MODELS, YUNQI_PROVIDER
|
||||
from apps.ai.model_routing import (
|
||||
ModelRequirements,
|
||||
match_model_requirements,
|
||||
model_metadata_errors,
|
||||
provider_metadata_errors,
|
||||
resolve_fallback_candidates,
|
||||
)
|
||||
from apps.ai.models import ModelConfig, ModelProvider
|
||||
|
||||
|
||||
def text_metadata(*, outbound=True, candidate=True):
|
||||
return {
|
||||
"routing": {
|
||||
"fallback_on_failure": outbound,
|
||||
"fallback_candidate": candidate,
|
||||
},
|
||||
"capabilities": {
|
||||
"operations": ["chat"],
|
||||
"features": ["streaming", "structured_output"],
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
class ModelRequirementsMatchTests(SimpleTestCase):
|
||||
def model(self, capability, metadata):
|
||||
return ModelConfig(name="match-model", display_name="Match", capability=capability, metadata=metadata)
|
||||
|
||||
def test_requirements_reject_invalid_reference_and_limit_values(self):
|
||||
with self.assertRaisesRegex(ValueError, "reference_images 不能为负数"):
|
||||
ModelRequirements(capability="image", operation="image_generate", reference_images=-1)
|
||||
with self.assertRaisesRegex(ValueError, "reference_images 必须等于 1"):
|
||||
ModelRequirements(
|
||||
capability="image",
|
||||
operation="image_edit",
|
||||
reference_mode="single",
|
||||
reference_images=2,
|
||||
)
|
||||
|
||||
def test_image_requires_exact_operation_reference_mode_count_and_ratio(self):
|
||||
model = self.model(
|
||||
ModelConfig.Capability.IMAGE,
|
||||
{
|
||||
"capabilities": {
|
||||
"operations": ["image_generate", "image_edit"],
|
||||
"features": [],
|
||||
"reference_modes": ["none", "single", "multiple"],
|
||||
"max_reference_images": 4,
|
||||
"aspect_ratios": ["1:1", "9:16"],
|
||||
}
|
||||
},
|
||||
)
|
||||
matched = match_model_requirements(
|
||||
model,
|
||||
ModelRequirements(
|
||||
capability="image",
|
||||
operation="image_edit",
|
||||
reference_mode="multiple",
|
||||
reference_images=4,
|
||||
aspect_ratio="9:16",
|
||||
),
|
||||
)
|
||||
rejected = match_model_requirements(
|
||||
model,
|
||||
ModelRequirements(
|
||||
capability="image",
|
||||
operation="image_edit",
|
||||
reference_mode="multiple",
|
||||
reference_images=5,
|
||||
aspect_ratio="16:9",
|
||||
),
|
||||
)
|
||||
|
||||
self.assertTrue(matched.matched)
|
||||
self.assertFalse(rejected.matched)
|
||||
self.assertTrue(any("参考图片上限不足" in reason for reason in rejected.reasons))
|
||||
self.assertTrue(any("画面比例" in reason for reason in rejected.reasons))
|
||||
|
||||
def test_video_requires_features_resolution_duration_and_exact_price_tier(self):
|
||||
model = self.model(
|
||||
ModelConfig.Capability.VIDEO,
|
||||
{
|
||||
"pricing": {"default": {"no_ref_video": 1}, "1080p": {"no_ref_video": 2}},
|
||||
"capabilities": {
|
||||
"operations": ["video_generate"],
|
||||
"features": ["text_to_video", "start_frame", "generate_audio"],
|
||||
"resolutions": ["720p", "1080p", "4k"],
|
||||
"durations": [5, 10],
|
||||
"aspect_ratios": ["16:9"],
|
||||
},
|
||||
},
|
||||
)
|
||||
matched = match_model_requirements(
|
||||
model,
|
||||
ModelRequirements(
|
||||
capability="video",
|
||||
operation="video_generate",
|
||||
features=frozenset({"start_frame", "generate_audio"}),
|
||||
resolution="1080p",
|
||||
duration=10,
|
||||
aspect_ratio="16:9",
|
||||
),
|
||||
)
|
||||
rejected = match_model_requirements(
|
||||
model,
|
||||
ModelRequirements(
|
||||
capability="video",
|
||||
operation="video_generate",
|
||||
features=frozenset({"video_reference"}),
|
||||
resolution="4k",
|
||||
duration=10,
|
||||
aspect_ratio="16:9",
|
||||
),
|
||||
)
|
||||
|
||||
self.assertTrue(matched.matched)
|
||||
self.assertFalse(rejected.matched)
|
||||
self.assertTrue(any("video_reference" in reason for reason in rejected.reasons))
|
||||
self.assertTrue(any("4k 的视频价格" in reason for reason in rejected.reasons))
|
||||
|
||||
def test_audio_requires_public_voice_mapping_and_limits(self):
|
||||
model = self.model(
|
||||
ModelConfig.Capability.AUDIO,
|
||||
{
|
||||
"capabilities": {
|
||||
"operations": ["tts"],
|
||||
"features": [],
|
||||
"languages": ["zh-CN"],
|
||||
"voice_map": {"narrator": "provider-voice-7"},
|
||||
"max_chars": 1000,
|
||||
"speed_range": [0.5, 2.0],
|
||||
"output_formats": ["mp3"],
|
||||
}
|
||||
},
|
||||
)
|
||||
matched = match_model_requirements(
|
||||
model,
|
||||
ModelRequirements(
|
||||
capability="audio",
|
||||
operation="tts",
|
||||
language="zh-CN",
|
||||
public_voice="narrator",
|
||||
char_count=999,
|
||||
speed_ratio=1.5,
|
||||
output_format="mp3",
|
||||
),
|
||||
)
|
||||
rejected = match_model_requirements(
|
||||
model,
|
||||
ModelRequirements(
|
||||
capability="audio",
|
||||
operation="tts",
|
||||
language="zh-CN",
|
||||
public_voice="unknown",
|
||||
char_count=1001,
|
||||
speed_ratio=3,
|
||||
output_format="wav",
|
||||
),
|
||||
)
|
||||
|
||||
self.assertTrue(matched.matched)
|
||||
self.assertFalse(rejected.matched)
|
||||
self.assertGreaterEqual(len(rejected.reasons), 4)
|
||||
|
||||
def test_missing_capability_contract_is_conservatively_rejected(self):
|
||||
model = self.model(ModelConfig.Capability.TEXT, {"routing": {"fallback_candidate": True}})
|
||||
result = match_model_requirements(model, ModelRequirements(capability="text", operation="chat"))
|
||||
|
||||
self.assertFalse(result.matched)
|
||||
self.assertIn("不支持操作 chat", result.reasons)
|
||||
|
||||
def test_chinese_metadata_validation_catches_bad_switches_and_priority(self):
|
||||
self.assertTrue(provider_metadata_errors({"routing": {"fallback_priority": True}}))
|
||||
errors = model_metadata_errors(
|
||||
"video",
|
||||
{
|
||||
"routing": {"fallback_on_failure": "yes", "fallback_candidate": True},
|
||||
"capabilities": {"operations": ["video_generate"]},
|
||||
},
|
||||
)
|
||||
|
||||
self.assertTrue(any("true 或 false" in error for error in errors))
|
||||
self.assertTrue(any("resolutions" in error for error in errors))
|
||||
self.assertTrue(any("durations" in error for error in errors))
|
||||
|
||||
def test_bootstrap_catalog_contains_persistent_routing_contracts(self):
|
||||
self.assertEqual(VOLCANO_PROVIDER["metadata"]["routing"]["fallback_priority"], 10)
|
||||
self.assertEqual(YUNQI_PROVIDER["metadata"]["routing"]["fallback_priority"], 20)
|
||||
for item in VOLCANO_MODELS + YUNQI_MODELS:
|
||||
self.assertIn("routing", item["metadata"])
|
||||
self.assertIn("capabilities", item["metadata"])
|
||||
self.assertTrue(item["metadata"]["routing"]["fallback_candidate"])
|
||||
self.assertTrue(item["metadata"]["capabilities"]["operations"])
|
||||
|
||||
|
||||
class DynamicFallbackCandidateTests(TestCase):
|
||||
requirements = ModelRequirements(capability="text", operation="chat")
|
||||
|
||||
def setUp(self):
|
||||
# 数据迁移会种入真实候选;路由单测只比较本用例创建的数据,避免环境顺序干扰。
|
||||
ModelConfig.objects.all().update(status=ModelConfig.Status.DISABLED)
|
||||
|
||||
def provider(self, name, priority=100, status=ModelProvider.Status.ACTIVE):
|
||||
return ModelProvider.objects.create(
|
||||
name=name,
|
||||
display_name=name,
|
||||
status=status,
|
||||
metadata={"routing": {"fallback_priority": priority}},
|
||||
)
|
||||
|
||||
def model(self, provider, name, *, outbound=True, candidate=True, status=ModelConfig.Status.ACTIVE, metadata=None):
|
||||
return ModelConfig.objects.create(
|
||||
provider=provider,
|
||||
name=name,
|
||||
display_name=name,
|
||||
capability=ModelConfig.Capability.TEXT,
|
||||
endpoint="chat/completions",
|
||||
status=status,
|
||||
metadata=metadata if metadata is not None else text_metadata(outbound=outbound, candidate=candidate),
|
||||
)
|
||||
|
||||
def policy(self, max_models=5):
|
||||
raw = deepcopy(settings.MODEL_ROUTING_POLICY)
|
||||
raw["max_models"] = max_models
|
||||
raw["max_calls"] = max(raw["max_calls"], max_models)
|
||||
return raw
|
||||
|
||||
def test_dynamic_order_is_priority_then_latest_updated_at(self):
|
||||
primary = self.model(self.provider("route-primary", 500), "primary")
|
||||
priority_10 = self.model(self.provider("route-priority-10", 10), "priority-10")
|
||||
priority_20 = self.model(self.provider("route-priority-20", 20), "priority-20")
|
||||
normal_provider = self.provider("route-normal", 100)
|
||||
normal_old = self.model(normal_provider, "normal-old")
|
||||
normal_new = self.model(normal_provider, "normal-new")
|
||||
now = timezone.now()
|
||||
ModelConfig.objects.filter(pk=normal_old.pk).update(updated_at=now - timedelta(days=1))
|
||||
ModelConfig.objects.filter(pk=normal_new.pk).update(updated_at=now)
|
||||
|
||||
with override_settings(MODEL_ROUTING_POLICY=self.policy()):
|
||||
candidates = resolve_fallback_candidates(primary_model=primary, requirements=self.requirements)
|
||||
|
||||
self.assertEqual(
|
||||
[model.name for model in candidates],
|
||||
[priority_10.name, priority_20.name, normal_new.name, normal_old.name],
|
||||
)
|
||||
|
||||
def test_user_primary_is_preserved_and_outbound_switch_can_stop_fallback(self):
|
||||
primary = self.model(self.provider("route-stop-primary", 500), "primary", outbound=False)
|
||||
self.model(self.provider("route-stop-candidate", 1), "candidate")
|
||||
|
||||
candidates = resolve_fallback_candidates(primary_model=primary, requirements=self.requirements)
|
||||
|
||||
self.assertEqual(candidates, [])
|
||||
self.assertEqual(primary.name, "primary")
|
||||
|
||||
def test_attempted_models_excluded_providers_and_disabled_rows_are_filtered(self):
|
||||
primary = self.model(self.provider("route-filter-primary", 500), "primary")
|
||||
attempted = self.model(self.provider("route-filter-attempted", 1), "attempted")
|
||||
excluded = self.model(self.provider("route-filter-provider", 2), "excluded")
|
||||
disabled_model = self.model(
|
||||
self.provider("route-filter-disabled-model", 3),
|
||||
"disabled-model",
|
||||
status=ModelConfig.Status.DISABLED,
|
||||
)
|
||||
disabled_provider = self.model(
|
||||
self.provider("route-filter-disabled-provider", 4, ModelProvider.Status.DISABLED),
|
||||
"disabled-provider",
|
||||
)
|
||||
valid = self.model(self.provider("route-filter-valid", 5), "valid")
|
||||
|
||||
with override_settings(MODEL_ROUTING_POLICY=self.policy()):
|
||||
candidates = resolve_fallback_candidates(
|
||||
primary_model=primary,
|
||||
requirements=self.requirements,
|
||||
attempted_model_ids={attempted.id},
|
||||
excluded_provider_ids={excluded.provider_id},
|
||||
)
|
||||
|
||||
self.assertEqual([model.id for model in candidates], [valid.id])
|
||||
self.assertNotIn(disabled_model.id, [model.id for model in candidates])
|
||||
self.assertNotIn(disabled_provider.id, [model.id for model in candidates])
|
||||
|
||||
def test_new_openai_compatible_provider_enters_pool_by_data_only(self):
|
||||
primary = self.model(self.provider("route-data-primary", 500), "primary")
|
||||
future_provider = self.provider("future-openai-compatible-provider", 15)
|
||||
future_model = self.model(future_provider, "future-text-model")
|
||||
|
||||
candidates = resolve_fallback_candidates(primary_model=primary, requirements=self.requirements)
|
||||
|
||||
self.assertEqual([model.id for model in candidates], [future_model.id])
|
||||
|
||||
def test_unconfigured_or_incompatible_candidates_are_conservatively_excluded(self):
|
||||
primary = self.model(self.provider("route-safe-primary", 500), "primary")
|
||||
self.model(self.provider("route-safe-missing", 1), "missing", metadata={})
|
||||
self.model(
|
||||
self.provider("route-safe-incompatible", 2),
|
||||
"incompatible",
|
||||
metadata={
|
||||
"routing": {"fallback_on_failure": True, "fallback_candidate": True},
|
||||
"capabilities": {"operations": ["embeddings"], "features": []},
|
||||
},
|
||||
)
|
||||
|
||||
candidates = resolve_fallback_candidates(primary_model=primary, requirements=self.requirements)
|
||||
|
||||
self.assertEqual(candidates, [])
|
||||
|
||||
def test_admin_serializers_reject_invalid_metadata_with_chinese_message(self):
|
||||
provider = self.provider("route-admin", 100)
|
||||
model = self.model(provider, "admin-model")
|
||||
provider_serializer = AdminModelProviderSerializer(
|
||||
provider,
|
||||
data={"metadata": {"routing": {"fallback_priority": "first"}}},
|
||||
partial=True,
|
||||
)
|
||||
model_serializer = AdminModelConfigSerializer(
|
||||
model,
|
||||
data={
|
||||
"metadata": {
|
||||
"routing": {"fallback_candidate": True, "fallback_on_failure": True},
|
||||
"capabilities": {},
|
||||
}
|
||||
},
|
||||
partial=True,
|
||||
)
|
||||
|
||||
self.assertFalse(provider_serializer.is_valid())
|
||||
self.assertIn("数字越小越优先", str(provider_serializer.errors))
|
||||
self.assertFalse(model_serializer.is_valid())
|
||||
self.assertIn("至少配置一个操作", str(model_serializer.errors))
|
||||
|
||||
def test_data_migration_seeded_existing_yunqi_rows(self):
|
||||
provider = ModelProvider.objects.filter(name="yunqi").first()
|
||||
model = ModelConfig.objects.filter(provider__name="yunqi", capability="image").first()
|
||||
|
||||
self.assertIsNotNone(provider)
|
||||
self.assertEqual(provider.metadata["routing"]["fallback_priority"], 20)
|
||||
self.assertIsNotNone(model)
|
||||
self.assertIn("routing", model.metadata)
|
||||
self.assertIn("capabilities", model.metadata)
|
||||
@@ -0,0 +1,216 @@
|
||||
from decimal import Decimal
|
||||
from io import BytesIO
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import requests
|
||||
from django.test import TestCase
|
||||
|
||||
from apps.accounts.models import Team, TeamMember, User
|
||||
from apps.ai.models import AITask, ModelConfig, ModelProvider
|
||||
from apps.ai.services import generate_model_triview, run_model_triview_task
|
||||
from apps.assets.models import Asset, AssetFile, Model
|
||||
from apps.billing.models import CreditAccount, CreditLedger
|
||||
|
||||
|
||||
def _metadata(*, outbound=True, base_cost="0.50"):
|
||||
return {
|
||||
"routing": {"fallback_on_failure": outbound, "fallback_candidate": True},
|
||||
"capabilities": {
|
||||
"operations": ["image_generate", "image_edit"],
|
||||
"features": [],
|
||||
"reference_modes": ["none", "single", "multiple"],
|
||||
"max_reference_images": 9,
|
||||
"aspect_ratios": ["1:1", "3:4", "4:5", "9:16", "16:9"],
|
||||
},
|
||||
"pricing": {"base_cost_yuan": base_cost},
|
||||
}
|
||||
|
||||
|
||||
class ModelTriviewRoutingTests(TestCase):
|
||||
def setUp(self):
|
||||
ModelConfig.objects.filter(capability=ModelConfig.Capability.IMAGE).update(
|
||||
status=ModelConfig.Status.DISABLED
|
||||
)
|
||||
self.user = User.objects.create_user(username="tri-routing", password="x")
|
||||
self.team = Team.objects.create(name="Tri Routing", owner=self.user)
|
||||
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
|
||||
CreditAccount.objects.create(team=self.team, balance=Decimal("1000"))
|
||||
self.portrait = Asset.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
name="模特形象",
|
||||
asset_type=Asset.Type.IMAGE,
|
||||
source=Asset.Source.UPLOAD,
|
||||
category=Asset.Category.MODEL_PORTRAIT,
|
||||
in_library=False,
|
||||
)
|
||||
AssetFile.objects.create(
|
||||
asset=self.portrait,
|
||||
object_key="portrait.png",
|
||||
bucket="bucket",
|
||||
content_type="image/png",
|
||||
preview_url="http://example.test/portrait.png",
|
||||
is_primary=True,
|
||||
)
|
||||
self.model_entity = Model.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
name="路由模特",
|
||||
portrait_asset=self.portrait,
|
||||
)
|
||||
self.provider_mocks = {}
|
||||
patch("apps.ai.services.get_image_provider", side_effect=self._provider_for).start()
|
||||
patch("apps.ai.tasks.generate_model_triview_task.delay").start()
|
||||
patch(
|
||||
"apps.ai.services.VolcanoArkProvider.media_to_bytes",
|
||||
return_value=(BytesIO(b"image"), "image/png"),
|
||||
).start()
|
||||
storage = patch("apps.ai.services.TosStorage").start()
|
||||
stored = storage.return_value.upload_fileobj.return_value
|
||||
stored.object_key = "tri.png"
|
||||
stored.bucket = "bucket"
|
||||
stored.content_type = "image/png"
|
||||
stored.size_bytes = 5
|
||||
patch("apps.assets.review.submit_asset_for_review").start()
|
||||
self.addCleanup(patch.stopall)
|
||||
|
||||
def provider(self, name, priority):
|
||||
return ModelProvider.objects.create(
|
||||
name=name,
|
||||
display_name=name,
|
||||
status=ModelProvider.Status.ACTIVE,
|
||||
metadata={"routing": {"fallback_priority": priority}},
|
||||
)
|
||||
|
||||
def model(self, provider, name, *, outbound=True, base_cost="0.50"):
|
||||
return ModelConfig.objects.create(
|
||||
provider=provider,
|
||||
name=name,
|
||||
display_name=name,
|
||||
capability=ModelConfig.Capability.IMAGE,
|
||||
endpoint="images/generations",
|
||||
unit_price=Decimal("20"),
|
||||
status=ModelConfig.Status.ACTIVE,
|
||||
metadata=_metadata(outbound=outbound, base_cost=base_cost),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _new_provider_mock():
|
||||
provider = Mock()
|
||||
provider.image_edit.return_value = {"data": [{"url": "http://example.test/tri.png"}]}
|
||||
provider.image_generation.return_value = {"data": [{"url": "http://example.test/tri.png"}]}
|
||||
provider.extract_first_media_url.side_effect = lambda response: response["data"][0]["url"]
|
||||
return provider
|
||||
|
||||
@staticmethod
|
||||
def _new_generation_only_provider_mock():
|
||||
provider = Mock(spec=["image_generation", "extract_first_media_url"])
|
||||
provider.image_generation.return_value = {"data": [{"url": "http://example.test/tri.png"}]}
|
||||
provider.extract_first_media_url.side_effect = lambda response: response["data"][0]["url"]
|
||||
return provider
|
||||
|
||||
def _provider_for(self, model):
|
||||
return self.provider_mocks.setdefault(model.id, self._new_provider_mock())
|
||||
|
||||
def submit(self):
|
||||
task, created = generate_model_triview(model=self.model_entity, user=self.user)
|
||||
self.assertTrue(created)
|
||||
return task
|
||||
|
||||
@staticmethod
|
||||
def ledger_count(task, ledger_type):
|
||||
return CreditLedger.objects.filter(task=task, ledger_type=ledger_type).count()
|
||||
|
||||
def test_first_success_records_single_reference_16_9_and_updates_model(self):
|
||||
primary = self.model(self.provider("tri-primary", 20), "tri-primary-model")
|
||||
task = self.submit()
|
||||
|
||||
run_model_triview_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
self.model_entity.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertTrue(task.request_payload["model_routing_v1"])
|
||||
attempt = task.model_attempts.get()
|
||||
self.assertEqual(attempt.model_config_id, primary.id)
|
||||
self.assertEqual(attempt.operation, "image_edit")
|
||||
self.assertEqual(attempt.request_summary["reference_images"], 1)
|
||||
self.assertEqual(attempt.request_summary["aspect_ratio"], "16:9")
|
||||
self.assertTrue(attempt.request_summary["model_triview"])
|
||||
call = self.provider_mocks[primary.id].image_edit.call_args
|
||||
self.assertEqual(call.kwargs["images"], ["http://example.test/portrait.png"])
|
||||
self.assertEqual(call.kwargs["size"], "1536x864")
|
||||
self.assertIsNotNone(self.model_entity.triview_asset_id)
|
||||
self.assertEqual(self.model_entity.triview_asset.origin_task_id, task.id)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
|
||||
def test_retry_then_dynamic_candidate_success_charges_once(self):
|
||||
primary = self.model(self.provider("tri-fallback-primary", 100), "tri-primary", base_cost="0.25")
|
||||
candidate = self.model(
|
||||
self.provider("volcano", 10), "tri-candidate", outbound=False, base_cost="0.75"
|
||||
)
|
||||
task = self.submit()
|
||||
self.provider_mocks[primary.id] = self._new_provider_mock()
|
||||
self.provider_mocks[primary.id].image_edit.side_effect = requests.ConnectionError("offline")
|
||||
|
||||
run_model_triview_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
attempts = list(task.model_attempts.all())
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertEqual([a.model_config_id for a in attempts], [primary.id, primary.id, candidate.id])
|
||||
self.assertTrue(attempts[-1].is_fallback)
|
||||
self.assertTrue(all(a.request_summary["aspect_ratio"] == "16:9" for a in attempts))
|
||||
self.assertTrue(all(a.public_model_name == "AirShelf Image" for a in attempts))
|
||||
self.assertEqual(task.base_cost, Decimal("0.7500"))
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
|
||||
|
||||
def test_all_candidates_fail_releases_and_keeps_existing_triview(self):
|
||||
primary = self.model(self.provider("tri-all-primary", 100), "tri-primary")
|
||||
candidate = self.model(self.provider("volcano", 10), "tri-candidate", outbound=False)
|
||||
old = Asset.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
name="旧三视图",
|
||||
asset_type=Asset.Type.IMAGE,
|
||||
source=Asset.Source.AI_GENERATED,
|
||||
category=Asset.Category.TRI_VIEW,
|
||||
in_library=False,
|
||||
)
|
||||
self.model_entity.triview_asset = old
|
||||
self.model_entity.save(update_fields=["triview_asset"])
|
||||
task = self.submit()
|
||||
for model in (primary, candidate):
|
||||
self.provider_mocks[model.id] = self._new_provider_mock()
|
||||
self.provider_mocks[model.id].image_edit.side_effect = requests.ConnectionError("offline")
|
||||
|
||||
run_model_triview_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
self.model_entity.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.FAILED)
|
||||
self.assertEqual(task.model_attempts.count(), 3)
|
||||
self.assertEqual(self.model_entity.triview_asset_id, old.id)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
|
||||
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
|
||||
|
||||
def test_direct_primary_uses_image_generation_without_outbound_fallback(self):
|
||||
primary = self.model(
|
||||
self.provider("volcano", 10), "seedream-triview", outbound=False
|
||||
)
|
||||
self.provider_mocks[primary.id] = self._new_generation_only_provider_mock()
|
||||
task = self.submit()
|
||||
|
||||
run_model_triview_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
call = self.provider_mocks[primary.id].image_generation.call_args
|
||||
self.assertEqual(call.kwargs["image"], ["http://example.test/portrait.png"])
|
||||
self.assertEqual(call.kwargs["size"], "1536x864")
|
||||
self.assertEqual(task.model_attempts.get().public_model_name, primary.display_name)
|
||||
@@ -0,0 +1,234 @@
|
||||
from decimal import Decimal
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import requests
|
||||
from django.test import TestCase
|
||||
|
||||
from apps.accounts.models import Team, TeamMember, User
|
||||
from apps.ai.models import AITask, ModelConfig, ModelProvider
|
||||
from apps.ai.services import generate_person_triview, run_triview_task
|
||||
from apps.assets.models import Asset, AssetFile, Model
|
||||
from apps.billing.models import CreditAccount, CreditLedger
|
||||
from apps.products.models import Product
|
||||
from apps.projects.models import BaseAssetGroup, Project
|
||||
|
||||
|
||||
def _metadata(*, outbound=True, base_cost="0.50"):
|
||||
return {
|
||||
"routing": {"fallback_on_failure": outbound, "fallback_candidate": True},
|
||||
"capabilities": {
|
||||
"operations": ["image_generate", "image_edit"],
|
||||
"features": [],
|
||||
"reference_modes": ["none", "single", "multiple"],
|
||||
"max_reference_images": 9,
|
||||
"aspect_ratios": ["1:1", "3:4", "4:5", "9:16", "16:9"],
|
||||
},
|
||||
"pricing": {"base_cost_yuan": base_cost},
|
||||
}
|
||||
|
||||
|
||||
class ProjectPersonTriviewRoutingTests(TestCase):
|
||||
def setUp(self):
|
||||
ModelConfig.objects.filter(capability=ModelConfig.Capability.IMAGE).update(
|
||||
status=ModelConfig.Status.DISABLED
|
||||
)
|
||||
self.user = User.objects.create_user(username="project-tri-routing", password="x")
|
||||
self.team = Team.objects.create(name="Project Tri Routing", owner=self.user)
|
||||
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
|
||||
CreditAccount.objects.create(team=self.team, balance=Decimal("1000"))
|
||||
product = Product.objects.create(team=self.team, created_by=self.user, title="测试商品")
|
||||
self.project = Project.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
product=product,
|
||||
name="项目角色三视图",
|
||||
)
|
||||
self.portrait = Asset.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
name="角色立绘",
|
||||
asset_type=Asset.Type.IMAGE,
|
||||
source=Asset.Source.AI_GENERATED,
|
||||
category=Asset.Category.PERSON,
|
||||
)
|
||||
AssetFile.objects.create(
|
||||
asset=self.portrait,
|
||||
object_key="portrait.png",
|
||||
bucket="bucket",
|
||||
content_type="image/png",
|
||||
preview_url="http://example.test/portrait.png",
|
||||
is_primary=True,
|
||||
)
|
||||
self.provider_mocks = {}
|
||||
patch("apps.ai.services.get_image_provider", side_effect=self._provider_for).start()
|
||||
patch("apps.ai.tasks.generate_triview_task.delay").start()
|
||||
patch("apps.ai.services._store_generated_media", side_effect=self._store_media).start()
|
||||
patch("apps.assets.review.submit_asset_for_review").start()
|
||||
self.addCleanup(patch.stopall)
|
||||
|
||||
def provider(self, name, priority):
|
||||
return ModelProvider.objects.create(
|
||||
name=name,
|
||||
display_name=name,
|
||||
status=ModelProvider.Status.ACTIVE,
|
||||
metadata={"routing": {"fallback_priority": priority}},
|
||||
)
|
||||
|
||||
def model(self, provider, name, *, outbound=True, base_cost="0.50"):
|
||||
return ModelConfig.objects.create(
|
||||
provider=provider,
|
||||
name=name,
|
||||
display_name=name,
|
||||
capability=ModelConfig.Capability.IMAGE,
|
||||
endpoint="images/generations",
|
||||
unit_price=Decimal("20"),
|
||||
status=ModelConfig.Status.ACTIVE,
|
||||
metadata=_metadata(outbound=outbound, base_cost=base_cost),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _new_provider_mock():
|
||||
provider = Mock()
|
||||
provider.image_edit.return_value = {"data": [{"url": "http://example.test/tri.png"}]}
|
||||
provider.image_generation.return_value = {"data": [{"url": "http://example.test/tri.png"}]}
|
||||
provider.extract_first_media_url.side_effect = lambda response: response["data"][0]["url"]
|
||||
return provider
|
||||
|
||||
@staticmethod
|
||||
def _new_generation_only_provider_mock():
|
||||
provider = Mock(spec=["image_generation", "extract_first_media_url"])
|
||||
provider.image_generation.return_value = {"data": [{"url": "http://example.test/tri.png"}]}
|
||||
provider.extract_first_media_url.side_effect = lambda response: response["data"][0]["url"]
|
||||
return provider
|
||||
|
||||
def _provider_for(self, model):
|
||||
return self.provider_mocks.setdefault(model.id, self._new_provider_mock())
|
||||
|
||||
def _store_media(self, **kwargs):
|
||||
return Asset.objects.create(
|
||||
team=kwargs["team"],
|
||||
created_by=kwargs["user"],
|
||||
name=kwargs["name"],
|
||||
asset_type=kwargs["asset_type"],
|
||||
source=Asset.Source.AI_GENERATED,
|
||||
category=kwargs["category"],
|
||||
origin_task=kwargs["task"],
|
||||
)
|
||||
|
||||
def submit(self):
|
||||
return generate_person_triview(
|
||||
project=self.project,
|
||||
user=self.user,
|
||||
portrait_asset=self.portrait,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def ledger_count(task, ledger_type):
|
||||
return CreditLedger.objects.filter(task=task, ledger_type=ledger_type).count()
|
||||
|
||||
def test_first_success_records_single_reference_16_9_and_keeps_project_scope(self):
|
||||
primary = self.model(self.provider("project-tri-primary", 20), "project-tri-primary")
|
||||
task = self.submit()
|
||||
|
||||
run_triview_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertTrue(task.request_payload["model_routing_v1"])
|
||||
attempt = task.model_attempts.get()
|
||||
self.assertEqual(attempt.model_config_id, primary.id)
|
||||
self.assertEqual(attempt.operation, "image_edit")
|
||||
self.assertEqual(attempt.request_summary["reference_images"], 1)
|
||||
self.assertEqual(attempt.request_summary["aspect_ratio"], "16:9")
|
||||
self.assertTrue(attempt.request_summary["project_person_triview"])
|
||||
call = self.provider_mocks[primary.id].image_edit.call_args
|
||||
self.assertEqual(call.kwargs["images"], ["http://example.test/portrait.png"])
|
||||
self.assertEqual(call.kwargs["size"], "1536x864")
|
||||
group = self.project.base_asset_groups.get(metadata__triview_of=str(self.portrait.id))
|
||||
self.assertEqual(group.adopted_asset.origin_task_id, task.id)
|
||||
self.assertFalse(Model.objects.filter(portrait_asset=self.portrait).exists())
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
|
||||
def test_retry_then_dynamic_candidate_success_charges_once(self):
|
||||
primary = self.model(
|
||||
self.provider("project-tri-fallback-primary", 100),
|
||||
"project-tri-primary",
|
||||
base_cost="0.25",
|
||||
)
|
||||
candidate = self.model(
|
||||
self.provider("volcano", 10),
|
||||
"project-tri-candidate",
|
||||
outbound=False,
|
||||
base_cost="0.75",
|
||||
)
|
||||
task = self.submit()
|
||||
self.provider_mocks[primary.id] = self._new_provider_mock()
|
||||
self.provider_mocks[primary.id].image_edit.side_effect = requests.ConnectionError("offline")
|
||||
|
||||
run_triview_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
attempts = list(task.model_attempts.all())
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertEqual([a.model_config_id for a in attempts], [primary.id, primary.id, candidate.id])
|
||||
self.assertTrue(attempts[-1].is_fallback)
|
||||
self.assertTrue(all(a.request_summary["aspect_ratio"] == "16:9" for a in attempts))
|
||||
self.assertEqual(task.base_cost, Decimal("0.7500"))
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
|
||||
|
||||
def test_all_candidates_fail_releases_and_keeps_existing_group(self):
|
||||
primary = self.model(self.provider("project-tri-all-primary", 100), "project-tri-primary")
|
||||
candidate = self.model(
|
||||
self.provider("volcano", 10), "project-tri-candidate", outbound=False
|
||||
)
|
||||
old = Asset.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
name="旧三视图",
|
||||
asset_type=Asset.Type.IMAGE,
|
||||
source=Asset.Source.AI_GENERATED,
|
||||
category=Asset.Category.TRI_VIEW,
|
||||
)
|
||||
group = BaseAssetGroup.objects.create(
|
||||
project=self.project,
|
||||
kind=BaseAssetGroup.Kind.PERSON,
|
||||
prompt="旧提示词",
|
||||
adopted_asset=old,
|
||||
metadata={"label": "·三视图", "triview_of": str(self.portrait.id)},
|
||||
)
|
||||
group.candidate_assets.add(old)
|
||||
task = self.submit()
|
||||
for model in (primary, candidate):
|
||||
self.provider_mocks[model.id] = self._new_provider_mock()
|
||||
self.provider_mocks[model.id].image_edit.side_effect = requests.ConnectionError("offline")
|
||||
|
||||
run_triview_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
group.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.FAILED)
|
||||
self.assertEqual(task.model_attempts.count(), 3)
|
||||
self.assertEqual(group.adopted_asset_id, old.id)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
|
||||
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
|
||||
|
||||
def test_direct_primary_uses_image_generation_with_reference(self):
|
||||
primary = self.model(
|
||||
self.provider("volcano", 10), "seedream-project-triview", outbound=False
|
||||
)
|
||||
self.provider_mocks[primary.id] = self._new_generation_only_provider_mock()
|
||||
task = self.submit()
|
||||
|
||||
run_triview_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
call = self.provider_mocks[primary.id].image_generation.call_args
|
||||
self.assertEqual(call.kwargs["image"], ["http://example.test/portrait.png"])
|
||||
self.assertEqual(call.kwargs["size"], "1536x864")
|
||||
self.assertEqual(task.model_attempts.get().public_model_name, primary.display_name)
|
||||
@@ -0,0 +1,366 @@
|
||||
from decimal import Decimal
|
||||
from unittest.mock import Mock
|
||||
|
||||
import requests
|
||||
from django.test import TestCase
|
||||
|
||||
from apps.accounts.models import Team, TeamMember, User
|
||||
from apps.ai.model_routing import ModelRequirements
|
||||
from apps.ai.models import AIModelAttempt, AITask, ModelConfig, ModelProvider
|
||||
from apps.ai.routing_executor import AttemptMetadata, execute_model_call, sanitize_attempt_summary
|
||||
from apps.billing.models import CreditAccount, CreditLedger, CreditReservation
|
||||
from apps.billing.services.ledger import charge_reserved_credit, release_credit, reserve_credit
|
||||
|
||||
|
||||
def text_metadata(*, outbound=True, candidate=True):
|
||||
return {
|
||||
"routing": {"fallback_on_failure": outbound, "fallback_candidate": candidate},
|
||||
"capabilities": {
|
||||
"operations": ["chat"],
|
||||
"features": ["streaming", "structured_output"],
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
class RoutingExecutorTests(TestCase):
|
||||
requirements = ModelRequirements(capability="text", operation="chat")
|
||||
|
||||
def setUp(self):
|
||||
# 迁移种子不会参与本组动态候选,保证断言只覆盖用例声明的模型。
|
||||
ModelConfig.objects.all().update(status=ModelConfig.Status.DISABLED)
|
||||
self.user = User.objects.create_user(username="routing-user", password="x")
|
||||
self.team = Team.objects.create(name="Routing Team", owner=self.user)
|
||||
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
|
||||
CreditAccount.objects.create(team=self.team, balance=Decimal("1000"))
|
||||
|
||||
def provider(self, name, priority):
|
||||
return ModelProvider.objects.create(
|
||||
name=name,
|
||||
display_name=name,
|
||||
metadata={"routing": {"fallback_priority": priority}},
|
||||
)
|
||||
|
||||
def model(self, provider, name, *, outbound=True, candidate=True):
|
||||
return ModelConfig.objects.create(
|
||||
provider=provider,
|
||||
name=name,
|
||||
display_name=name,
|
||||
capability=ModelConfig.Capability.TEXT,
|
||||
endpoint="chat/completions",
|
||||
metadata=text_metadata(outbound=outbound, candidate=candidate),
|
||||
)
|
||||
|
||||
def task(self, model, key):
|
||||
return AITask.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
model_config=model,
|
||||
task_type=AITask.Type.SCRIPT_GENERATION,
|
||||
status=AITask.Status.RESERVED,
|
||||
idempotency_key=key,
|
||||
estimated_cost=Decimal("10"),
|
||||
actual_cost=Decimal("10"),
|
||||
)
|
||||
|
||||
def execute(self, *, task, primary, invoke, **kwargs):
|
||||
return execute_model_call(
|
||||
task=task,
|
||||
primary_model=primary,
|
||||
requirements=self.requirements,
|
||||
public_model_name="AirShelf Script",
|
||||
invoke=invoke,
|
||||
sleep=lambda seconds: None,
|
||||
uniform=lambda start, end: 0,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def test_first_success_logs_once_and_never_queries_candidates(self):
|
||||
primary = self.model(self.provider("exec-first", 100), "primary")
|
||||
task = self.task(primary, "exec-first")
|
||||
resolver = Mock(return_value=[])
|
||||
|
||||
result = self.execute(
|
||||
task=task,
|
||||
primary=primary,
|
||||
invoke=lambda model, timeout: {"ok": True},
|
||||
candidate_resolver=resolver,
|
||||
)
|
||||
|
||||
self.assertEqual(result.actual_model, primary)
|
||||
self.assertFalse(result.fallback_used)
|
||||
self.assertEqual(result.call_count, 1)
|
||||
resolver.assert_not_called()
|
||||
attempt = task.model_attempts.get()
|
||||
self.assertEqual(attempt.status, AIModelAttempt.Status.SUCCEEDED)
|
||||
self.assertFalse(attempt.is_retry)
|
||||
self.assertFalse(attempt.is_fallback)
|
||||
|
||||
def test_transient_failure_retries_primary_then_succeeds(self):
|
||||
primary = self.model(self.provider("exec-retry", 100), "primary")
|
||||
task = self.task(primary, "exec-retry")
|
||||
resolver = Mock(return_value=[])
|
||||
calls = []
|
||||
|
||||
def invoke(model, timeout):
|
||||
calls.append((model.id, timeout))
|
||||
if len(calls) == 1:
|
||||
raise requests.ConnectionError("temporary network failure")
|
||||
return {"ok": True}
|
||||
|
||||
result = self.execute(
|
||||
task=task,
|
||||
primary=primary,
|
||||
invoke=invoke,
|
||||
candidate_resolver=resolver,
|
||||
)
|
||||
|
||||
self.assertEqual(result.call_count, 2)
|
||||
resolver.assert_not_called()
|
||||
attempts = list(task.model_attempts.all())
|
||||
self.assertEqual([item.status for item in attempts], ["failed", "succeeded"])
|
||||
self.assertEqual([item.is_retry for item in attempts], [False, True])
|
||||
self.assertEqual(attempts[1].previous_attempt_id, attempts[0].id)
|
||||
|
||||
def test_non_retryable_provider_failure_switches_and_preserves_primary(self):
|
||||
primary = self.model(self.provider("exec-fallback-primary", 100), "primary")
|
||||
candidate = self.model(self.provider("exec-fallback-candidate", 10), "candidate")
|
||||
task = self.task(primary, "exec-fallback")
|
||||
resolver = Mock(return_value=[candidate])
|
||||
|
||||
def invoke(model, timeout):
|
||||
if model.id == primary.id:
|
||||
raise RuntimeError("API_KEY is not configured")
|
||||
return {"ok": True}
|
||||
|
||||
result = self.execute(
|
||||
task=task,
|
||||
primary=primary,
|
||||
invoke=invoke,
|
||||
candidate_resolver=resolver,
|
||||
)
|
||||
|
||||
self.assertEqual(result.actual_model, candidate)
|
||||
self.assertTrue(result.fallback_used)
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(task.model_config_id, primary.id)
|
||||
attempts = list(task.model_attempts.all())
|
||||
self.assertEqual([item.model_name for item in attempts], ["primary", "candidate"])
|
||||
self.assertEqual([item.public_model_name for item in attempts], ["AirShelf Script", "AirShelf Script"])
|
||||
self.assertTrue(attempts[1].is_fallback)
|
||||
excluded = resolver.call_args.kwargs["excluded_provider_ids"]
|
||||
self.assertIn(primary.provider_id, excluded)
|
||||
|
||||
def test_all_models_fail_without_cycles_and_stops_at_global_call_limit(self):
|
||||
primary = self.model(self.provider("exec-all-primary", 500), "primary")
|
||||
self.model(self.provider("exec-all-first", 10), "candidate-1")
|
||||
self.model(self.provider("exec-all-second", 20), "candidate-2")
|
||||
task = self.task(primary, "exec-all")
|
||||
|
||||
with self.assertRaises(requests.ConnectionError):
|
||||
self.execute(
|
||||
task=task,
|
||||
primary=primary,
|
||||
invoke=lambda model, timeout: (_ for _ in ()).throw(requests.ConnectionError("offline")),
|
||||
)
|
||||
|
||||
attempts = list(task.model_attempts.all())
|
||||
self.assertEqual(len(attempts), 5)
|
||||
self.assertEqual([item.model_name for item in attempts], ["primary", "primary", "primary", "candidate-1", "candidate-2"])
|
||||
self.assertEqual(len({item.sequence for item in attempts}), 5)
|
||||
self.assertTrue(all(item.status == AIModelAttempt.Status.FAILED for item in attempts))
|
||||
|
||||
def test_content_rejection_neither_retries_nor_fallbacks(self):
|
||||
primary = self.model(self.provider("exec-safety", 100), "primary")
|
||||
task = self.task(primary, "exec-safety")
|
||||
resolver = Mock(return_value=[])
|
||||
|
||||
with self.assertRaises(RuntimeError):
|
||||
self.execute(
|
||||
task=task,
|
||||
primary=primary,
|
||||
invoke=lambda model, timeout: (_ for _ in ()).throw(RuntimeError("moderation_blocked")),
|
||||
candidate_resolver=resolver,
|
||||
)
|
||||
|
||||
resolver.assert_not_called()
|
||||
attempt = task.model_attempts.get()
|
||||
self.assertEqual(attempt.error_type, "content_rejected")
|
||||
|
||||
def test_retry_after_above_cap_skips_wait_and_switches_immediately(self):
|
||||
primary = self.model(self.provider("exec-rate-primary", 100), "primary")
|
||||
candidate = self.model(self.provider("exec-rate-candidate", 10), "candidate")
|
||||
task = self.task(primary, "exec-rate")
|
||||
sleeps = []
|
||||
|
||||
def invoke(model, timeout):
|
||||
if model.id == primary.id:
|
||||
response = requests.Response()
|
||||
response.status_code = 429
|
||||
response.headers["Retry-After"] = "60"
|
||||
exc = requests.HTTPError("rate limited", response=response)
|
||||
raise exc
|
||||
return {"ok": True}
|
||||
|
||||
result = execute_model_call(
|
||||
task=task,
|
||||
primary_model=primary,
|
||||
requirements=self.requirements,
|
||||
public_model_name="AirShelf Script",
|
||||
invoke=invoke,
|
||||
candidate_resolver=Mock(return_value=[candidate]),
|
||||
sleep=sleeps.append,
|
||||
uniform=lambda start, end: 0,
|
||||
)
|
||||
|
||||
self.assertEqual(result.call_count, 2)
|
||||
self.assertEqual(sleeps, [])
|
||||
|
||||
def test_video_provider_task_id_on_error_never_retries_or_fallbacks(self):
|
||||
provider = self.provider("exec-video", 100)
|
||||
model = ModelConfig.objects.create(
|
||||
provider=provider,
|
||||
name="video-model",
|
||||
display_name="video-model",
|
||||
capability=ModelConfig.Capability.VIDEO,
|
||||
endpoint="contents/generations/tasks",
|
||||
metadata={
|
||||
"routing": {"fallback_on_failure": True, "fallback_candidate": True},
|
||||
"capabilities": {
|
||||
"operations": ["video_generate"],
|
||||
"features": ["text_to_video"],
|
||||
"resolutions": ["720p"],
|
||||
"durations": [5],
|
||||
},
|
||||
"pricing": {"default": {"no_ref_video": 1}},
|
||||
},
|
||||
)
|
||||
task = self.task(model, "exec-video-known-id")
|
||||
resolver = Mock(return_value=[])
|
||||
|
||||
with self.assertRaises(requests.Timeout):
|
||||
execute_model_call(
|
||||
task=task,
|
||||
primary_model=model,
|
||||
requirements=ModelRequirements(
|
||||
capability="video",
|
||||
operation="video_generate",
|
||||
resolution="720p",
|
||||
duration=5,
|
||||
),
|
||||
public_model_name="Seedance",
|
||||
invoke=lambda current, timeout: (_ for _ in ()).throw(requests.Timeout("poll state unknown")),
|
||||
error_metadata=lambda exc, current: AttemptMetadata(provider_task_id="provider-task-123"),
|
||||
candidate_resolver=resolver,
|
||||
sleep=lambda seconds: None,
|
||||
)
|
||||
|
||||
resolver.assert_not_called()
|
||||
self.assertEqual(task.model_attempts.count(), 1)
|
||||
self.assertEqual(task.model_attempts.get().provider_task_id, "provider-task-123")
|
||||
|
||||
def test_metadata_extraction_failure_does_not_repeat_successful_model_call(self):
|
||||
primary = self.model(self.provider("exec-meta", 100), "primary")
|
||||
task = self.task(primary, "exec-meta")
|
||||
invoke = Mock(return_value={"ok": True})
|
||||
|
||||
result = self.execute(
|
||||
task=task,
|
||||
primary=primary,
|
||||
invoke=invoke,
|
||||
result_metadata=lambda value, model: (_ for _ in ()).throw(RuntimeError("summary failed")),
|
||||
)
|
||||
|
||||
self.assertEqual(result.call_count, 1)
|
||||
invoke.assert_called_once()
|
||||
self.assertIn("metadata_error", task.model_attempts.get().response_summary)
|
||||
|
||||
def test_attempt_costs_are_audited_and_summed_without_touching_parent_updated_at(self):
|
||||
primary = self.model(self.provider("exec-cost-primary", 100), "primary")
|
||||
candidate = self.model(self.provider("exec-cost-candidate", 10), "candidate")
|
||||
task = self.task(primary, "exec-cost")
|
||||
original_updated_at = task.updated_at
|
||||
|
||||
def invoke(model, timeout):
|
||||
if model.id == primary.id:
|
||||
exc = RuntimeError("API_KEY is not configured")
|
||||
exc.attempt_metadata = {"platform_cost": "0.125", "usage": {"prompt_tokens": 10}}
|
||||
raise exc
|
||||
return {"ok": True}
|
||||
|
||||
self.execute(
|
||||
task=task,
|
||||
primary=primary,
|
||||
invoke=invoke,
|
||||
candidate_resolver=Mock(return_value=[candidate]),
|
||||
result_metadata=lambda result, model: AttemptMetadata(
|
||||
platform_cost=Decimal("0.375"), usage={"output_tokens": 20}
|
||||
),
|
||||
)
|
||||
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(task.base_cost, Decimal("0.5000"))
|
||||
self.assertEqual(task.updated_at, original_updated_at)
|
||||
self.assertEqual(
|
||||
list(task.model_attempts.values_list("platform_cost", flat=True)),
|
||||
[Decimal("0.1250"), Decimal("0.3750")],
|
||||
)
|
||||
|
||||
def test_feedback_summary_redacts_credentials_and_large_binary(self):
|
||||
summary = sanitize_attempt_summary(
|
||||
{
|
||||
"api_key": "secret-value",
|
||||
"authorization": "Bearer secret",
|
||||
"prompt": "hello",
|
||||
"image": b"123456",
|
||||
}
|
||||
)
|
||||
|
||||
self.assertEqual(summary["api_key"], "[REDACTED]")
|
||||
self.assertEqual(summary["authorization"], "[REDACTED]")
|
||||
self.assertEqual(summary["image"], "[二进制 6 字节]")
|
||||
|
||||
def test_fallback_success_keeps_single_reserve_and_single_charge(self):
|
||||
primary = self.model(self.provider("exec-bill-primary", 100), "primary")
|
||||
candidate = self.model(self.provider("exec-bill-candidate", 10), "candidate")
|
||||
task = self.task(primary, "exec-bill-success")
|
||||
reservation = reserve_credit(team=self.team, user=self.user, task=task, amount=Decimal("10"))
|
||||
|
||||
def invoke(model, timeout):
|
||||
if model.id == primary.id:
|
||||
raise RuntimeError("API_KEY is not configured")
|
||||
return {"ok": True}
|
||||
|
||||
self.execute(
|
||||
task=task,
|
||||
primary=primary,
|
||||
invoke=invoke,
|
||||
candidate_resolver=Mock(return_value=[candidate]),
|
||||
)
|
||||
charge_reserved_credit(reservation=reservation, actual_amount=Decimal("10"))
|
||||
|
||||
self.assertEqual(CreditReservation.objects.filter(task=task).count(), 1)
|
||||
self.assertEqual(CreditLedger.objects.filter(task=task, ledger_type=CreditLedger.Type.RESERVE).count(), 1)
|
||||
self.assertEqual(CreditLedger.objects.filter(task=task, ledger_type=CreditLedger.Type.CHARGE).count(), 1)
|
||||
self.assertEqual(task.model_attempts.count(), 2)
|
||||
|
||||
def test_all_failures_keep_single_reserve_and_single_release(self):
|
||||
primary = self.model(self.provider("exec-release-primary", 500), "primary")
|
||||
self.model(self.provider("exec-release-first", 10), "candidate-1")
|
||||
self.model(self.provider("exec-release-second", 20), "candidate-2")
|
||||
task = self.task(primary, "exec-bill-failure")
|
||||
reservation = reserve_credit(team=self.team, user=self.user, task=task, amount=Decimal("10"))
|
||||
|
||||
with self.assertRaises(requests.ConnectionError):
|
||||
self.execute(
|
||||
task=task,
|
||||
primary=primary,
|
||||
invoke=lambda model, timeout: (_ for _ in ()).throw(requests.ConnectionError("offline")),
|
||||
)
|
||||
release_credit(reservation=reservation, reason="all providers failed")
|
||||
|
||||
self.assertEqual(CreditReservation.objects.filter(task=task).count(), 1)
|
||||
self.assertEqual(CreditLedger.objects.filter(task=task, ledger_type=CreditLedger.Type.RESERVE).count(), 1)
|
||||
self.assertEqual(CreditLedger.objects.filter(task=task, ledger_type=CreditLedger.Type.CHARGE).count(), 0)
|
||||
self.assertEqual(CreditLedger.objects.filter(task=task, ledger_type=CreditLedger.Type.RELEASE).count(), 1)
|
||||
account = CreditAccount.objects.get(team=self.team)
|
||||
self.assertEqual(account.reserved_balance, Decimal("0"))
|
||||
@@ -0,0 +1,157 @@
|
||||
from copy import deepcopy
|
||||
from dataclasses import FrozenInstanceError
|
||||
import os
|
||||
from unittest.mock import patch
|
||||
|
||||
from django.apps import apps
|
||||
from django.conf import settings
|
||||
from django.core.exceptions import ImproperlyConfigured
|
||||
from django.test import SimpleTestCase, override_settings
|
||||
|
||||
from apps.ai.routing_policy import RoutingPolicyConfigurationError, load_model_routing_policy
|
||||
from airshelf.settings.base import env_float, env_int, env_int_list
|
||||
|
||||
|
||||
class ModelRoutingPolicyTests(SimpleTestCase):
|
||||
def policy_dict(self) -> dict:
|
||||
return deepcopy(settings.MODEL_ROUTING_POLICY)
|
||||
|
||||
def test_default_policy_matches_confirmed_values_and_is_immutable(self):
|
||||
policy = load_model_routing_policy()
|
||||
|
||||
self.assertEqual(policy.max_models, 3)
|
||||
self.assertEqual(policy.max_calls, 5)
|
||||
self.assertEqual(policy.jitter_ratio, 0.20)
|
||||
self.assertEqual(policy.text.retry_delays, (1.0, 3.0))
|
||||
self.assertEqual(policy.text.request_timeout, 120.0)
|
||||
self.assertEqual(policy.text.stream_timeout, 300.0)
|
||||
self.assertEqual(policy.text.total_timeout, 480.0)
|
||||
self.assertEqual(policy.image.retry_delays, (3.0,))
|
||||
self.assertEqual(policy.image.request_timeout, 300.0)
|
||||
self.assertEqual(policy.image.total_timeout, 900.0)
|
||||
self.assertEqual(policy.audio.retry_delays, (2.0,))
|
||||
self.assertEqual(policy.audio.total_timeout, 180.0)
|
||||
self.assertEqual(policy.video.submit_retry_delays, (3.0,))
|
||||
self.assertEqual(policy.video.submit_timeout, 120.0)
|
||||
self.assertEqual(policy.video.submit_total_timeout, 300.0)
|
||||
self.assertEqual(policy.video.poll_request_timeout, 60.0)
|
||||
self.assertEqual(policy.video.generation_timeout, 1800.0)
|
||||
self.assertEqual(policy.postprocess.retry_delays, (2.0, 5.0, 10.0))
|
||||
|
||||
with self.assertRaises(FrozenInstanceError):
|
||||
policy.max_calls = 6
|
||||
|
||||
def test_settings_can_inject_short_timeouts_without_production_branch(self):
|
||||
raw = self.policy_dict()
|
||||
raw["text"].update(
|
||||
retry_delays=[0],
|
||||
request_timeout=1,
|
||||
stream_timeout=1,
|
||||
total_timeout=2,
|
||||
retry_after_cap=0,
|
||||
)
|
||||
raw["image"].update(retry_delays=[], request_timeout=1, total_timeout=1, retry_after_cap=0)
|
||||
raw["audio"].update(retry_delays=[], request_timeout=1, total_timeout=1, retry_after_cap=0)
|
||||
raw["video"].update(
|
||||
submit_retry_delays=[],
|
||||
submit_timeout=1,
|
||||
submit_total_timeout=1,
|
||||
poll_request_timeout=1,
|
||||
generation_timeout=1,
|
||||
retry_after_cap=0,
|
||||
)
|
||||
raw["postprocess"]["retry_delays"] = []
|
||||
|
||||
with override_settings(MODEL_ROUTING_POLICY=raw):
|
||||
policy = load_model_routing_policy()
|
||||
|
||||
self.assertEqual(policy.text.retry_delays, (0.0,))
|
||||
self.assertEqual(policy.text.total_timeout, 2.0)
|
||||
self.assertEqual(policy.image.retry_delays, ())
|
||||
self.assertEqual(policy.video.generation_timeout, 1.0)
|
||||
self.assertEqual(policy.postprocess.retry_delays, ())
|
||||
|
||||
def test_rejects_max_calls_smaller_than_model_count_with_chinese_message(self):
|
||||
raw = self.policy_dict()
|
||||
raw["max_models"] = 3
|
||||
raw["max_calls"] = 2
|
||||
|
||||
with self.assertRaisesRegex(
|
||||
RoutingPolicyConfigurationError,
|
||||
"最多真实模型调用次数.*不得小于 max_models.*最多尝试模型数量",
|
||||
):
|
||||
load_model_routing_policy(raw)
|
||||
|
||||
def test_rejects_single_timeout_greater_than_total_timeout(self):
|
||||
raw = self.policy_dict()
|
||||
raw["image"]["request_timeout"] = 901
|
||||
raw["image"]["total_timeout"] = 900
|
||||
|
||||
with self.assertRaisesRegex(
|
||||
RoutingPolicyConfigurationError,
|
||||
"图片|request_timeout.*单次调用超时.*不得大于.*total_timeout",
|
||||
):
|
||||
load_model_routing_policy(raw)
|
||||
|
||||
def test_rejects_non_numeric_retry_delay_and_reports_value(self):
|
||||
raw = self.policy_dict()
|
||||
raw["audio"]["retry_delays"] = ["两秒"]
|
||||
|
||||
with self.assertRaisesRegex(
|
||||
RoutingPolicyConfigurationError,
|
||||
r"audio\.retry_delays\[0\].*必须是数字.*两秒",
|
||||
):
|
||||
load_model_routing_policy(raw)
|
||||
|
||||
def test_rejects_missing_required_section(self):
|
||||
raw = self.policy_dict()
|
||||
del raw["postprocess"]
|
||||
|
||||
with self.assertRaisesRegex(
|
||||
RoutingPolicyConfigurationError,
|
||||
"MODEL_ROUTING_POLICY.postprocess.*必填配置",
|
||||
):
|
||||
load_model_routing_policy(raw)
|
||||
|
||||
def test_rejects_video_poll_timeout_longer_than_generation_window(self):
|
||||
raw = self.policy_dict()
|
||||
raw["video"]["poll_request_timeout"] = 61
|
||||
raw["video"]["generation_timeout"] = 60
|
||||
|
||||
with self.assertRaisesRegex(
|
||||
RoutingPolicyConfigurationError,
|
||||
"视频单次轮询超时.*不得大于.*generation_timeout",
|
||||
):
|
||||
load_model_routing_policy(raw)
|
||||
|
||||
def test_app_startup_fails_fast_for_invalid_policy(self):
|
||||
raw = self.policy_dict()
|
||||
raw["jitter_ratio"] = 1.5
|
||||
|
||||
with override_settings(MODEL_ROUTING_POLICY=raw):
|
||||
with self.assertRaisesRegex(ImproperlyConfigured, "模型路由策略配置错误.*重试随机抖动比例"):
|
||||
apps.get_app_config("ai").ready()
|
||||
|
||||
|
||||
class RoutingPolicyEnvironmentParserTests(SimpleTestCase):
|
||||
def test_integer_float_and_list_environment_overrides(self):
|
||||
with patch.dict(
|
||||
os.environ,
|
||||
{
|
||||
"TEST_ROUTING_INT": "7",
|
||||
"TEST_ROUTING_FLOAT": "0.25",
|
||||
"TEST_ROUTING_LIST": "1, 3,5",
|
||||
},
|
||||
):
|
||||
self.assertEqual(env_int("TEST_ROUTING_INT", 1), 7)
|
||||
self.assertEqual(env_float("TEST_ROUTING_FLOAT", 0.1), 0.25)
|
||||
self.assertEqual(env_int_list("TEST_ROUTING_LIST", (9,)), [1, 3, 5])
|
||||
|
||||
def test_empty_retry_delay_environment_means_no_retry(self):
|
||||
with patch.dict(os.environ, {"TEST_ROUTING_LIST": ""}):
|
||||
self.assertEqual(env_int_list("TEST_ROUTING_LIST", (1, 3)), [])
|
||||
|
||||
def test_invalid_environment_value_has_chinese_error(self):
|
||||
with patch.dict(os.environ, {"TEST_ROUTING_INT": "三"}):
|
||||
with self.assertRaisesRegex(ValueError, "环境变量 TEST_ROUTING_INT 必须是整数.*三"):
|
||||
env_int("TEST_ROUTING_INT", 1)
|
||||
@@ -0,0 +1,179 @@
|
||||
from decimal import Decimal
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
from django.test import TestCase
|
||||
|
||||
from apps.accounts.models import Team, TeamMember, User
|
||||
from apps.ai.models import AITask, ModelConfig, ModelProvider
|
||||
from apps.ai.script_agent import persist_script_draft
|
||||
from apps.ai.services import extract_cast_and_scenes
|
||||
from apps.billing.models import CreditAccount, CreditLedger
|
||||
from apps.products.models import Product
|
||||
from apps.projects.models import Project
|
||||
|
||||
|
||||
class ScriptEntitySyncTests(TestCase):
|
||||
def setUp(self):
|
||||
ModelConfig.objects.filter(capability=ModelConfig.Capability.TEXT).update(
|
||||
status=ModelConfig.Status.DISABLED
|
||||
)
|
||||
self.user = User.objects.create_user(username="script-entity-sync", password="x")
|
||||
self.team = Team.objects.create(name="Script Entity Sync", owner=self.user)
|
||||
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
|
||||
CreditAccount.objects.create(team=self.team, balance=Decimal("1000"))
|
||||
self.product = Product.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
title="测试商品",
|
||||
)
|
||||
self.project = Project.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
product=self.product,
|
||||
name="脚本实体同步项目",
|
||||
)
|
||||
provider = ModelProvider.objects.create(
|
||||
name="script-sync-provider",
|
||||
display_name="Script Sync Provider",
|
||||
status=ModelProvider.Status.ACTIVE,
|
||||
)
|
||||
self.model = ModelConfig.objects.create(
|
||||
provider=provider,
|
||||
name="script-sync-model",
|
||||
display_name="Script Sync Model",
|
||||
capability=ModelConfig.Capability.TEXT,
|
||||
endpoint="chat/completions",
|
||||
unit_price=Decimal("10"),
|
||||
status=ModelConfig.Status.ACTIVE,
|
||||
metadata={
|
||||
"routing": {"fallback_on_failure": True, "fallback_candidate": True},
|
||||
"capabilities": {
|
||||
"operations": ["chat"],
|
||||
"features": ["structured_output"],
|
||||
},
|
||||
"pricing": {"base_cost_yuan": "0.50"},
|
||||
},
|
||||
)
|
||||
|
||||
def main_task(self):
|
||||
return AITask.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
project=self.project,
|
||||
task_type=AITask.Type.SCRIPT_GENERATION,
|
||||
status=AITask.Status.SUBMITTED,
|
||||
model_config=self.model,
|
||||
idempotency_key=f"script-entity-sync:{self.project.id}",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def draft():
|
||||
return {
|
||||
"hook": "开场钩子",
|
||||
"tone": "自然",
|
||||
"aspect_ratio": "9:16",
|
||||
"total_duration": 15,
|
||||
"segment_count": 1,
|
||||
"entities": [
|
||||
{
|
||||
"id": "c1",
|
||||
"type": "character",
|
||||
"name": "女主",
|
||||
"visual_prompt": "都市女主",
|
||||
"ref_index": 1,
|
||||
},
|
||||
{
|
||||
"id": "s1",
|
||||
"type": "scene",
|
||||
"name": "客厅",
|
||||
"visual_prompt": "现代客厅",
|
||||
"ref_index": 2,
|
||||
},
|
||||
],
|
||||
"segments": [
|
||||
{
|
||||
"index": 0,
|
||||
"duration": 15,
|
||||
"narration": "女主在客厅展示商品",
|
||||
"visual": "女主站在客厅",
|
||||
"role": "钩子",
|
||||
"speaker": "女主",
|
||||
"product_exposure": "展示",
|
||||
"entity_refs": ["c1", "s1"],
|
||||
"dialogue": [],
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
def test_script_draft_syncs_entities_without_creating_second_ai_task(self):
|
||||
task = self.main_task()
|
||||
before_count = AITask.objects.count()
|
||||
|
||||
script = persist_script_draft(
|
||||
project=self.project,
|
||||
user=self.user,
|
||||
task=task,
|
||||
draft=self.draft(),
|
||||
source="ai",
|
||||
)
|
||||
|
||||
self.project.refresh_from_db()
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(AITask.objects.count(), before_count)
|
||||
self.assertEqual(task.status, AITask.Status.SUBMITTED)
|
||||
self.assertEqual(script.task_id, task.id)
|
||||
self.assertEqual(self.project.metadata["cast"], ["女主"])
|
||||
self.assertEqual(self.project.metadata["scenes"], ["客厅"])
|
||||
self.assertEqual(
|
||||
[entity["id"] for entity in self.project.metadata["script_entities"]],
|
||||
["c1", "s1"],
|
||||
)
|
||||
self.assertEqual(CreditLedger.objects.filter(task=task).count(), 0)
|
||||
|
||||
@patch("apps.ai.services.get_text_provider")
|
||||
def test_legacy_post_script_extractor_uses_routed_attempt_and_one_settlement(self, get_provider):
|
||||
provider = Mock()
|
||||
response = {
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"content": (
|
||||
'{"cast":[{"name":"女主","prompt":"都市女主"}],'
|
||||
'"scenes":[{"name":"客厅","prompt":"现代客厅"}]}'
|
||||
)
|
||||
}
|
||||
}
|
||||
],
|
||||
"usage": {"total_tokens": 42},
|
||||
}
|
||||
provider.chat_completion.return_value = response
|
||||
provider.extract_text.return_value = response["choices"][0]["message"]["content"]
|
||||
get_provider.return_value = provider
|
||||
|
||||
result = extract_cast_and_scenes(
|
||||
project=self.project,
|
||||
user=self.user,
|
||||
content="女主在客厅展示商品",
|
||||
)
|
||||
|
||||
task = AITask.objects.get(task_type=AITask.Type.SCRIPT_OPTIMIZATION)
|
||||
self.assertEqual(result["cast"], ["女主"])
|
||||
self.assertEqual(result["scenes"], ["客厅"])
|
||||
self.assertTrue(task.request_payload["model_routing_v1"])
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
attempt = task.model_attempts.get()
|
||||
self.assertEqual(attempt.operation, "chat")
|
||||
self.assertEqual(attempt.public_model_name, "AirShelf Script")
|
||||
self.assertEqual(attempt.request_summary["source"], "legacy_post_script_extract")
|
||||
self.assertEqual(
|
||||
CreditLedger.objects.filter(task=task, ledger_type=CreditLedger.Type.RESERVE).count(),
|
||||
1,
|
||||
)
|
||||
self.assertEqual(
|
||||
CreditLedger.objects.filter(task=task, ledger_type=CreditLedger.Type.CHARGE).count(),
|
||||
1,
|
||||
)
|
||||
self.assertEqual(
|
||||
CreditLedger.objects.filter(task=task, ledger_type=CreditLedger.Type.RELEASE).count(),
|
||||
0,
|
||||
)
|
||||
@@ -0,0 +1,226 @@
|
||||
import json
|
||||
from decimal import Decimal
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import requests
|
||||
from django.test import TestCase
|
||||
|
||||
from apps.accounts.models import Team, TeamMember, User
|
||||
from apps.ai.models import AITask, ModelConfig, ModelProvider
|
||||
from apps.ai.script_agent import regenerate_segment_via_agent
|
||||
from apps.billing.models import CreditAccount, CreditLedger
|
||||
from apps.products.models import Product
|
||||
from apps.projects.models import Project, ScriptSegment, ScriptVersion
|
||||
|
||||
|
||||
def _metadata(*, outbound=True, base_cost="0.50"):
|
||||
return {
|
||||
"routing": {"fallback_on_failure": outbound, "fallback_candidate": True},
|
||||
"capabilities": {
|
||||
"operations": ["chat"],
|
||||
"features": ["structured_output"],
|
||||
},
|
||||
"pricing": {"base_cost_yuan": base_cost},
|
||||
}
|
||||
|
||||
|
||||
class ScriptOptimizationRoutingTests(TestCase):
|
||||
def setUp(self):
|
||||
ModelConfig.objects.filter(capability=ModelConfig.Capability.TEXT).update(
|
||||
status=ModelConfig.Status.DISABLED
|
||||
)
|
||||
self.user = User.objects.create_user(username="script-opt-routing", password="x")
|
||||
self.team = Team.objects.create(name="Script Opt Routing", owner=self.user)
|
||||
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
|
||||
CreditAccount.objects.create(team=self.team, balance=Decimal("1000"))
|
||||
product = Product.objects.create(team=self.team, created_by=self.user, title="测试商品")
|
||||
self.project = Project.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
product=product,
|
||||
name="单镜优化路由项目",
|
||||
)
|
||||
self.base_draft = {
|
||||
"hook": "钩子",
|
||||
"tone": "自然",
|
||||
"aspect_ratio": "9:16",
|
||||
"total_duration": 30,
|
||||
"segment_count": 2,
|
||||
"entities": [],
|
||||
"segments": [
|
||||
{
|
||||
"index": 0,
|
||||
"duration": 15,
|
||||
"role": "钩子",
|
||||
"narration": "旧口播0",
|
||||
"visual": "旧画面0",
|
||||
"speaker": None,
|
||||
"product_exposure": "",
|
||||
"entity_refs": [],
|
||||
"dialogue": [],
|
||||
},
|
||||
{
|
||||
"index": 1,
|
||||
"duration": 15,
|
||||
"role": "CTA",
|
||||
"narration": "旧口播1",
|
||||
"visual": "旧画面1",
|
||||
"speaker": None,
|
||||
"product_exposure": "",
|
||||
"entity_refs": [],
|
||||
"dialogue": [],
|
||||
},
|
||||
],
|
||||
}
|
||||
self.script = ScriptVersion.objects.create(
|
||||
project=self.project,
|
||||
title="基准稿",
|
||||
content=json.dumps(self.base_draft, ensure_ascii=False),
|
||||
is_adopted=True,
|
||||
metadata={
|
||||
"hook": "钩子",
|
||||
"tone": "自然",
|
||||
"aspect_ratio": "9:16",
|
||||
"total_duration": 30,
|
||||
"entities": [],
|
||||
},
|
||||
)
|
||||
ScriptSegment.objects.create(
|
||||
script_version=self.script,
|
||||
sort_order=0,
|
||||
duration_seconds=15,
|
||||
role="钩子",
|
||||
narration="旧口播0",
|
||||
visual_prompt="旧画面0",
|
||||
)
|
||||
self.target = ScriptSegment.objects.create(
|
||||
script_version=self.script,
|
||||
sort_order=1,
|
||||
duration_seconds=15,
|
||||
role="CTA",
|
||||
narration="旧口播1",
|
||||
visual_prompt="旧画面1",
|
||||
)
|
||||
self.provider_mocks = {}
|
||||
patch("apps.ai.services.get_text_provider", side_effect=self._provider_for).start()
|
||||
self.addCleanup(patch.stopall)
|
||||
|
||||
def provider(self, name, priority):
|
||||
return ModelProvider.objects.create(
|
||||
name=name,
|
||||
display_name=name,
|
||||
status=ModelProvider.Status.ACTIVE,
|
||||
metadata={"routing": {"fallback_priority": priority}},
|
||||
)
|
||||
|
||||
def model(self, provider, name, *, outbound=True, base_cost="0.50"):
|
||||
return ModelConfig.objects.create(
|
||||
provider=provider,
|
||||
name=name,
|
||||
display_name=name,
|
||||
capability=ModelConfig.Capability.TEXT,
|
||||
endpoint="chat/completions",
|
||||
unit_price=Decimal("10"),
|
||||
status=ModelConfig.Status.ACTIVE,
|
||||
metadata=_metadata(outbound=outbound, base_cost=base_cost),
|
||||
)
|
||||
|
||||
def valid_text(self):
|
||||
draft = json.loads(json.dumps(self.base_draft))
|
||||
draft["segments"][1]["narration"] = "优化后的口播"
|
||||
draft["segments"][1]["visual"] = "优化后的画面"
|
||||
return json.dumps(draft, ensure_ascii=False)
|
||||
|
||||
def good_provider(self):
|
||||
provider = Mock()
|
||||
provider.chat_completion.return_value = {"usage": {"total_tokens": 42}}
|
||||
provider.extract_text.return_value = self.valid_text()
|
||||
return provider
|
||||
|
||||
def _provider_for(self, model):
|
||||
return self.provider_mocks.setdefault(model.id, self.good_provider())
|
||||
|
||||
@staticmethod
|
||||
def ledger_count(task, ledger_type):
|
||||
return CreditLedger.objects.filter(task=task, ledger_type=ledger_type).count()
|
||||
|
||||
def optimize(self, primary):
|
||||
return regenerate_segment_via_agent(
|
||||
project=self.project,
|
||||
user=self.user,
|
||||
model_config=primary,
|
||||
segment=self.target,
|
||||
instruction="让第二镜更有吸引力",
|
||||
)
|
||||
|
||||
def test_first_success_records_attempt_and_preserves_other_segment(self):
|
||||
primary = self.model(self.provider("script-opt-primary", 20), "script-opt-primary")
|
||||
|
||||
version = self.optimize(primary)
|
||||
|
||||
task = AITask.objects.get(task_type=AITask.Type.SCRIPT_OPTIMIZATION)
|
||||
segments = list(version.segments.order_by("sort_order"))
|
||||
self.assertEqual(segments[0].narration, "旧口播0")
|
||||
self.assertEqual(segments[1].narration, "优化后的口播")
|
||||
self.assertTrue(task.request_payload["model_routing_v1"])
|
||||
attempt = task.model_attempts.get()
|
||||
self.assertEqual(attempt.model_config_id, primary.id)
|
||||
self.assertEqual(attempt.operation, "chat")
|
||||
self.assertEqual(attempt.request_summary["target_index"], 1)
|
||||
self.assertEqual(attempt.public_model_name, "AirShelf Script")
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
|
||||
def test_primary_retries_then_dynamic_fallback_succeeds_once(self):
|
||||
primary = self.model(
|
||||
self.provider("script-opt-fallback-primary", 100),
|
||||
"script-opt-primary",
|
||||
base_cost="0.25",
|
||||
)
|
||||
candidate = self.model(
|
||||
self.provider("script-opt-fallback-candidate", 10),
|
||||
"script-opt-candidate",
|
||||
outbound=False,
|
||||
base_cost="0.75",
|
||||
)
|
||||
failed = self.good_provider()
|
||||
failed.chat_completion.side_effect = requests.ConnectionError("offline")
|
||||
self.provider_mocks[primary.id] = failed
|
||||
|
||||
version = self.optimize(primary)
|
||||
|
||||
task = AITask.objects.get(task_type=AITask.Type.SCRIPT_OPTIMIZATION)
|
||||
attempts = list(task.model_attempts.all())
|
||||
self.assertEqual(
|
||||
[attempt.model_config_id for attempt in attempts],
|
||||
[primary.id, primary.id, primary.id, candidate.id],
|
||||
)
|
||||
self.assertTrue(attempts[-1].is_fallback)
|
||||
self.assertEqual(version.segments.order_by("sort_order")[0].narration, "旧口播0")
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
|
||||
|
||||
def test_all_fail_releases_once_and_does_not_create_new_version(self):
|
||||
primary = self.model(self.provider("script-opt-all-primary", 100), "script-opt-primary")
|
||||
fallback = self.model(
|
||||
self.provider("script-opt-all-fallback", 10),
|
||||
"script-opt-fallback",
|
||||
outbound=False,
|
||||
)
|
||||
for model in (primary, fallback):
|
||||
failed = self.good_provider()
|
||||
failed.chat_completion.side_effect = requests.ConnectionError("offline")
|
||||
self.provider_mocks[model.id] = failed
|
||||
before_versions = ScriptVersion.objects.filter(project=self.project).count()
|
||||
|
||||
with self.assertRaises(Exception):
|
||||
self.optimize(primary)
|
||||
|
||||
task = AITask.objects.get(task_type=AITask.Type.SCRIPT_OPTIMIZATION)
|
||||
self.assertEqual(task.status, AITask.Status.FAILED)
|
||||
self.assertEqual(ScriptVersion.objects.filter(project=self.project).count(), before_versions)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
|
||||
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
|
||||
@@ -0,0 +1,349 @@
|
||||
import json
|
||||
import time
|
||||
from decimal import Decimal
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import requests
|
||||
from django.test import TransactionTestCase
|
||||
|
||||
from apps.accounts.models import Team, TeamMember, User
|
||||
from apps.ai.models import AITask, ModelConfig, ModelProvider
|
||||
from apps.ai.script_agent import stream_script_agent
|
||||
from apps.billing.models import CreditAccount, CreditLedger
|
||||
from apps.products.models import Product
|
||||
from apps.projects.models import Project, ScriptSegment, ScriptVersion
|
||||
|
||||
|
||||
def _metadata(*, outbound=True, base_cost="0.50"):
|
||||
return {
|
||||
"routing": {"fallback_on_failure": outbound, "fallback_candidate": True},
|
||||
"capabilities": {
|
||||
"operations": ["chat"],
|
||||
"features": ["streaming", "structured_output"],
|
||||
},
|
||||
"pricing": {"base_cost_yuan": base_cost},
|
||||
}
|
||||
|
||||
|
||||
class ScriptStreamRoutingTests(TransactionTestCase):
|
||||
reset_sequences = True
|
||||
|
||||
def setUp(self):
|
||||
ModelConfig.objects.filter(capability=ModelConfig.Capability.TEXT).update(
|
||||
status=ModelConfig.Status.DISABLED
|
||||
)
|
||||
self.user = User.objects.create_user(username="script-stream-routing", password="x")
|
||||
self.team = Team.objects.create(name="Script Stream Routing", owner=self.user)
|
||||
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
|
||||
CreditAccount.objects.create(team=self.team, balance=Decimal("1000"))
|
||||
product = Product.objects.create(team=self.team, created_by=self.user, title="测试商品")
|
||||
self.project = Project.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
product=product,
|
||||
name="流式脚本路由项目",
|
||||
)
|
||||
self.provider_mocks = {}
|
||||
patch("apps.ai.services.get_text_provider", side_effect=self._provider_for).start()
|
||||
self.addCleanup(patch.stopall)
|
||||
|
||||
def provider(self, name, priority):
|
||||
return ModelProvider.objects.create(
|
||||
name=name,
|
||||
display_name=name,
|
||||
status=ModelProvider.Status.ACTIVE,
|
||||
metadata={"routing": {"fallback_priority": priority}},
|
||||
)
|
||||
|
||||
def model(self, provider, name, *, outbound=True, base_cost="0.50"):
|
||||
return ModelConfig.objects.create(
|
||||
provider=provider,
|
||||
name=name,
|
||||
display_name=name,
|
||||
capability=ModelConfig.Capability.TEXT,
|
||||
endpoint="chat/completions",
|
||||
unit_price=Decimal("10"),
|
||||
status=ModelConfig.Status.ACTIVE,
|
||||
metadata=_metadata(outbound=outbound, base_cost=base_cost),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def valid_draft():
|
||||
return {
|
||||
"hook": "开场钩子",
|
||||
"tone": "自然",
|
||||
"aspect_ratio": "9:16",
|
||||
"total_duration": 15,
|
||||
"segment_count": 1,
|
||||
"entities": [
|
||||
{
|
||||
"id": "c1",
|
||||
"type": "character",
|
||||
"name": "女主",
|
||||
"visual_prompt": "都市女主",
|
||||
"ref_index": 1,
|
||||
}
|
||||
],
|
||||
"segments": [
|
||||
{
|
||||
"index": 0,
|
||||
"duration": 15,
|
||||
"role": "钩子",
|
||||
"narration": "全新脚本口播",
|
||||
"visual": "女主展示商品",
|
||||
"speaker": "女主",
|
||||
"product_exposure": "展示",
|
||||
"entity_refs": ["c1"],
|
||||
"dialogue": [],
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
def good_provider(self):
|
||||
provider = Mock()
|
||||
raw = json.dumps(self.valid_draft(), ensure_ascii=False)
|
||||
provider.chat_completion_stream.return_value = iter(
|
||||
[
|
||||
{"type": "reasoning", "text": "先分析商品卖点"},
|
||||
{"type": "delta", "text": "我先给你一版。\n"},
|
||||
{"type": "delta", "text": raw},
|
||||
{"type": "done"},
|
||||
]
|
||||
)
|
||||
return provider
|
||||
|
||||
def _provider_for(self, model):
|
||||
return self.provider_mocks.setdefault(model.id, self.good_provider())
|
||||
|
||||
@staticmethod
|
||||
def parse_events(frames):
|
||||
return [json.loads(frame.removeprefix("data: ").strip()) for frame in frames]
|
||||
|
||||
@staticmethod
|
||||
def ledger_count(task, ledger_type):
|
||||
return CreditLedger.objects.filter(task=task, ledger_type=ledger_type).count()
|
||||
|
||||
def generate(self, primary):
|
||||
return self.parse_events(
|
||||
list(
|
||||
stream_script_agent(
|
||||
project=self.project,
|
||||
user=self.user,
|
||||
model_config=primary,
|
||||
mode="auto",
|
||||
user_prompt="生成一版脚本",
|
||||
aspect_ratio="9:16",
|
||||
total_duration=15,
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
def test_first_success_keeps_reasoning_and_delta_sse_and_records_attempt(self):
|
||||
primary = self.model(self.provider("script-stream-primary", 20), "script-stream-primary")
|
||||
|
||||
events = self.generate(primary)
|
||||
|
||||
task = AITask.objects.get(task_type=AITask.Type.SCRIPT_GENERATION)
|
||||
self.assertIn("reasoning", [event["type"] for event in events])
|
||||
self.assertIn("delta", [event["type"] for event in events])
|
||||
self.assertIn("draft", [event["type"] for event in events])
|
||||
self.assertEqual(events[-1]["type"], "done")
|
||||
self.assertEqual(ScriptVersion.objects.filter(project=self.project).count(), 1)
|
||||
attempt = task.model_attempts.get()
|
||||
self.assertEqual(attempt.model_config_id, primary.id)
|
||||
self.assertTrue(attempt.request_summary["streaming"])
|
||||
self.assertTrue(attempt.request_summary["structured_output"])
|
||||
self.assertEqual(attempt.request_summary["business_operation"], "script_generate")
|
||||
self.assertEqual(attempt.public_model_name, "AirShelf Script")
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
|
||||
def test_primary_retries_then_fallback_succeeds_with_one_script_and_settlement(self):
|
||||
primary = self.model(
|
||||
self.provider("script-stream-fallback-primary", 100),
|
||||
"script-stream-primary",
|
||||
base_cost="0.25",
|
||||
)
|
||||
candidate = self.model(
|
||||
self.provider("script-stream-fallback-candidate", 10),
|
||||
"script-stream-candidate",
|
||||
outbound=False,
|
||||
base_cost="0.75",
|
||||
)
|
||||
failed = Mock()
|
||||
failed.chat_completion_stream.side_effect = requests.ConnectionError("offline")
|
||||
self.provider_mocks[primary.id] = failed
|
||||
|
||||
events = self.generate(primary)
|
||||
|
||||
task = AITask.objects.get(task_type=AITask.Type.SCRIPT_GENERATION)
|
||||
attempts = list(task.model_attempts.all())
|
||||
self.assertEqual(events[-1]["type"], "done")
|
||||
self.assertEqual(ScriptVersion.objects.filter(project=self.project).count(), 1)
|
||||
self.assertEqual(
|
||||
[attempt.model_config_id for attempt in attempts],
|
||||
[primary.id, primary.id, primary.id, candidate.id],
|
||||
)
|
||||
self.assertTrue(attempts[-1].is_fallback)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
|
||||
|
||||
def test_all_fail_emits_safe_error_releases_once_and_saves_no_script(self):
|
||||
primary = self.model(self.provider("script-stream-all-primary", 100), "script-stream-primary")
|
||||
fallback = self.model(
|
||||
self.provider("script-stream-all-fallback", 10),
|
||||
"script-stream-fallback",
|
||||
outbound=False,
|
||||
)
|
||||
for model in (primary, fallback):
|
||||
failed = Mock()
|
||||
failed.chat_completion_stream.side_effect = requests.ConnectionError("offline")
|
||||
self.provider_mocks[model.id] = failed
|
||||
|
||||
events = self.generate(primary)
|
||||
|
||||
task = AITask.objects.get(task_type=AITask.Type.SCRIPT_GENERATION)
|
||||
self.assertEqual(task.status, AITask.Status.FAILED)
|
||||
self.assertEqual(events[-1]["type"], "error")
|
||||
self.assertNotIn("offline", events[-1]["detail"])
|
||||
self.assertEqual(ScriptVersion.objects.filter(project=self.project).count(), 0)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
|
||||
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
|
||||
|
||||
def test_client_disconnect_stops_stream_and_releases_without_saving_script(self):
|
||||
primary = self.model(self.provider("script-stream-abort", 20), "script-stream-primary")
|
||||
raw = json.dumps(self.valid_draft(), ensure_ascii=False)
|
||||
|
||||
def delayed_events():
|
||||
yield {"type": "reasoning", "text": "正在分析"}
|
||||
time.sleep(0.05)
|
||||
yield {"type": "delta", "text": raw}
|
||||
yield {"type": "done"}
|
||||
|
||||
provider = Mock()
|
||||
provider.chat_completion_stream.return_value = delayed_events()
|
||||
self.provider_mocks[primary.id] = provider
|
||||
stream = stream_script_agent(
|
||||
project=self.project,
|
||||
user=self.user,
|
||||
model_config=primary,
|
||||
mode="auto",
|
||||
user_prompt="生成一版脚本",
|
||||
aspect_ratio="9:16",
|
||||
total_duration=15,
|
||||
)
|
||||
for frame in stream:
|
||||
event = json.loads(frame.removeprefix("data: ").strip())
|
||||
if event["type"] == "reasoning":
|
||||
stream.close()
|
||||
break
|
||||
|
||||
deadline = time.monotonic() + 1
|
||||
task = AITask.objects.get(task_type=AITask.Type.SCRIPT_GENERATION)
|
||||
while not task.model_attempts.filter(status="failed").exists() and time.monotonic() < deadline:
|
||||
time.sleep(0.01)
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.FAILED)
|
||||
self.assertEqual(ScriptVersion.objects.filter(project=self.project).count(), 0)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
|
||||
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
|
||||
|
||||
def test_stream_revise_target_keeps_other_segment_and_uses_base_duration(self):
|
||||
primary = self.model(self.provider("script-stream-revise", 20), "script-stream-primary")
|
||||
base_draft = {
|
||||
"hook": "原钩子",
|
||||
"tone": "自然",
|
||||
"aspect_ratio": "9:16",
|
||||
"total_duration": 30,
|
||||
"segment_count": 2,
|
||||
"entities": [],
|
||||
"segments": [
|
||||
{
|
||||
"index": 0,
|
||||
"duration": 15,
|
||||
"role": "钩子",
|
||||
"narration": "旧口播0",
|
||||
"visual": "旧画面0",
|
||||
"speaker": None,
|
||||
"product_exposure": "",
|
||||
"entity_refs": [],
|
||||
"dialogue": [],
|
||||
},
|
||||
{
|
||||
"index": 1,
|
||||
"duration": 15,
|
||||
"role": "CTA",
|
||||
"narration": "旧口播1",
|
||||
"visual": "旧画面1",
|
||||
"speaker": None,
|
||||
"product_exposure": "",
|
||||
"entity_refs": [],
|
||||
"dialogue": [],
|
||||
},
|
||||
],
|
||||
}
|
||||
base = ScriptVersion.objects.create(
|
||||
project=self.project,
|
||||
title="基准稿",
|
||||
content=json.dumps(base_draft, ensure_ascii=False),
|
||||
metadata={"hook": "原钩子", "entities": [], "total_duration": 30},
|
||||
)
|
||||
ScriptSegment.objects.create(
|
||||
script_version=base,
|
||||
sort_order=0,
|
||||
duration_seconds=15,
|
||||
narration="旧口播0",
|
||||
visual_prompt="旧画面0",
|
||||
)
|
||||
ScriptSegment.objects.create(
|
||||
script_version=base,
|
||||
sort_order=1,
|
||||
duration_seconds=15,
|
||||
narration="旧口播1",
|
||||
visual_prompt="旧画面1",
|
||||
)
|
||||
revised = json.loads(json.dumps(base_draft))
|
||||
revised["segments"][1]["narration"] = "流式优化口播"
|
||||
revised["segments"][1]["visual"] = "流式优化画面"
|
||||
provider = Mock()
|
||||
provider.chat_completion_stream.return_value = iter(
|
||||
[
|
||||
{"type": "delta", "text": json.dumps(revised, ensure_ascii=False)},
|
||||
{"type": "done"},
|
||||
]
|
||||
)
|
||||
self.provider_mocks[primary.id] = provider
|
||||
|
||||
events = self.parse_events(
|
||||
list(
|
||||
stream_script_agent(
|
||||
project=self.project,
|
||||
user=self.user,
|
||||
model_config=primary,
|
||||
mode="revise",
|
||||
user_prompt="只改第二镜",
|
||||
base_version_id=str(base.id),
|
||||
aspect_ratio="9:16",
|
||||
total_duration=60,
|
||||
target_index=1,
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
task = AITask.objects.get(task_type=AITask.Type.SCRIPT_OPTIMIZATION)
|
||||
saved = ScriptVersion.objects.filter(project=self.project).order_by("-created_at").first()
|
||||
segments = list(saved.segments.order_by("sort_order"))
|
||||
self.assertEqual(events[-1]["type"], "done")
|
||||
self.assertEqual(len(segments), 2)
|
||||
self.assertEqual(segments[0].narration, "旧口播0")
|
||||
self.assertEqual(segments[1].narration, "流式优化口播")
|
||||
attempt = task.model_attempts.get()
|
||||
self.assertEqual(attempt.request_summary["target_index"], 1)
|
||||
self.assertEqual(attempt.request_summary["total_duration"], 30)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
@@ -0,0 +1,696 @@
|
||||
from decimal import Decimal
|
||||
from io import BytesIO
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import requests
|
||||
from django.test import TestCase
|
||||
|
||||
from apps.accounts.models import Team, TeamMember, User
|
||||
from apps.ai.models import AIModelAttempt, AITask, ModelConfig, ModelProvider
|
||||
from apps.ai.services import enqueue_standalone_images, run_standalone_image_task
|
||||
from apps.assets.models import Asset, AssetFile
|
||||
from apps.billing.models import CreditAccount, CreditLedger
|
||||
from apps.products.models import Product
|
||||
|
||||
|
||||
def image_metadata(
|
||||
*,
|
||||
outbound,
|
||||
candidate=True,
|
||||
ratios=None,
|
||||
base_cost="0.50",
|
||||
reference_modes=None,
|
||||
max_reference_images=0,
|
||||
):
|
||||
modes = reference_modes or ["none"]
|
||||
operations = ["image_generate"]
|
||||
if set(modes) & {"single", "multiple"}:
|
||||
operations.append("image_edit")
|
||||
return {
|
||||
"routing": {"fallback_on_failure": outbound, "fallback_candidate": candidate},
|
||||
"capabilities": {
|
||||
"operations": operations,
|
||||
"features": [],
|
||||
"reference_modes": modes,
|
||||
"max_reference_images": max_reference_images,
|
||||
"aspect_ratios": ratios or ["1:1", "3:4", "4:5", "9:16", "16:9"],
|
||||
},
|
||||
"pricing": {"base_cost_yuan": base_cost},
|
||||
}
|
||||
|
||||
|
||||
class StandaloneSingleImageRoutingTests(TestCase):
|
||||
def setUp(self):
|
||||
ModelConfig.objects.filter(capability=ModelConfig.Capability.IMAGE).update(status=ModelConfig.Status.DISABLED)
|
||||
self.user = User.objects.create_user(username="standalone-routing", password="x")
|
||||
self.team = Team.objects.create(name="Standalone Routing", owner=self.user)
|
||||
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
|
||||
CreditAccount.objects.create(team=self.team, balance=Decimal("1000"))
|
||||
self.provider_mocks = {}
|
||||
|
||||
provider_patch = patch("apps.ai.services.get_image_provider", side_effect=self._provider_for)
|
||||
provider_patch.start()
|
||||
media_patch = patch(
|
||||
"apps.ai.services.VolcanoArkProvider.media_to_bytes",
|
||||
return_value=(BytesIO(b"image"), "image/png"),
|
||||
)
|
||||
self.media_to_bytes = media_patch.start()
|
||||
storage_patch = patch("apps.ai.services.TosStorage")
|
||||
storage = storage_patch.start()
|
||||
stored = storage.return_value.upload_fileobj.return_value
|
||||
stored.object_key = "generated.png"
|
||||
stored.bucket = "bucket"
|
||||
stored.content_type = "image/png"
|
||||
stored.size_bytes = 5
|
||||
self.addCleanup(patch.stopall)
|
||||
|
||||
def provider(self, name, priority):
|
||||
provider, _ = ModelProvider.objects.update_or_create(
|
||||
name=name,
|
||||
defaults={
|
||||
"display_name": name,
|
||||
"status": ModelProvider.Status.ACTIVE,
|
||||
"metadata": {"routing": {"fallback_priority": priority}},
|
||||
},
|
||||
)
|
||||
return provider
|
||||
|
||||
def model(
|
||||
self,
|
||||
provider,
|
||||
name,
|
||||
*,
|
||||
outbound=True,
|
||||
ratios=None,
|
||||
base_cost="0.50",
|
||||
reference_modes=None,
|
||||
max_reference_images=0,
|
||||
):
|
||||
return ModelConfig.objects.create(
|
||||
provider=provider,
|
||||
name=name,
|
||||
display_name=name,
|
||||
capability=ModelConfig.Capability.IMAGE,
|
||||
endpoint="images/generations",
|
||||
unit_price=Decimal("20"),
|
||||
status=ModelConfig.Status.ACTIVE,
|
||||
metadata=image_metadata(
|
||||
outbound=outbound,
|
||||
ratios=ratios,
|
||||
base_cost=base_cost,
|
||||
reference_modes=reference_modes,
|
||||
max_reference_images=max_reference_images,
|
||||
),
|
||||
)
|
||||
|
||||
def _provider_for(self, model):
|
||||
return self.provider_mocks.setdefault(model.id, self._new_provider_mock())
|
||||
|
||||
@staticmethod
|
||||
def _new_provider_mock():
|
||||
provider = Mock()
|
||||
provider.image_generation.return_value = {"data": [{"url": "http://example.test/result.png"}]}
|
||||
provider.image_edit.return_value = {"data": [{"url": "http://example.test/result.png"}]}
|
||||
provider.extract_first_media_url.side_effect = lambda response: response["data"][0]["url"]
|
||||
return provider
|
||||
|
||||
@staticmethod
|
||||
def _new_generation_only_provider_mock():
|
||||
provider = Mock(spec=["image_generation", "extract_first_media_url"])
|
||||
provider.image_generation.return_value = {"data": [{"url": "http://example.test/result.png"}]}
|
||||
provider.extract_first_media_url.side_effect = lambda response: response["data"][0]["url"]
|
||||
return provider
|
||||
|
||||
def reference_asset(self, name):
|
||||
asset = Asset.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
name=name,
|
||||
asset_type=Asset.Type.IMAGE,
|
||||
source=Asset.Source.UPLOAD,
|
||||
category=Asset.Category.UPLOAD,
|
||||
)
|
||||
AssetFile.objects.create(
|
||||
asset=asset,
|
||||
object_key=f"{name}.png",
|
||||
bucket="bucket",
|
||||
content_type="image/png",
|
||||
preview_url=f"http://example.test/{name}.png",
|
||||
is_primary=True,
|
||||
)
|
||||
return asset
|
||||
|
||||
def submit(self, primary, *, count=1, ratio="1:1", reference_image_ids=None):
|
||||
task = enqueue_standalone_images(
|
||||
team=self.team,
|
||||
user=self.user,
|
||||
prompt="一只在太空漂浮的橙色猫",
|
||||
mode="image",
|
||||
count=count,
|
||||
ratio=ratio,
|
||||
image_model=f"{primary.provider.name}:{primary.name}",
|
||||
reference_image_ids=reference_image_ids,
|
||||
dispatch=False,
|
||||
)
|
||||
return task
|
||||
|
||||
def tryon_assets(self):
|
||||
product = Product.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
title="橙色童装上衣",
|
||||
category="服饰内衣",
|
||||
)
|
||||
cover = self.reference_asset("tryon-product")
|
||||
cover.category = Asset.Category.PRODUCT_IMAGE
|
||||
cover.save(update_fields=["category"])
|
||||
product.cover_asset = cover
|
||||
product.save(update_fields=["cover_asset"])
|
||||
model_asset = self.reference_asset("tryon-model")
|
||||
model_asset.source = Asset.Source.AI_GENERATED
|
||||
model_asset.category = Asset.Category.PERSON
|
||||
model_asset.save(update_fields=["source", "category"])
|
||||
return product, model_asset
|
||||
|
||||
def submit_tryon(self, primary, *, count=1, ratio="4:5"):
|
||||
product, model_asset = self.tryon_assets()
|
||||
tasks = enqueue_standalone_images(
|
||||
team=self.team,
|
||||
user=self.user,
|
||||
prompt="保持商品和模特身份,生成自然站姿上身图",
|
||||
mode="model",
|
||||
count=count,
|
||||
product_id=str(product.id),
|
||||
model_id=str(model_asset.id),
|
||||
ratio=ratio,
|
||||
image_model=f"{primary.provider.name}:{primary.name}",
|
||||
dispatch=False,
|
||||
)
|
||||
return tasks
|
||||
|
||||
def platform_product(self):
|
||||
product = Product.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
title="橙色童装上衣",
|
||||
category="服饰内衣",
|
||||
)
|
||||
cover = self.reference_asset("platform-product")
|
||||
cover.category = Asset.Category.PRODUCT_IMAGE
|
||||
cover.save(update_fields=["category"])
|
||||
product.cover_asset = cover
|
||||
product.save(update_fields=["cover_asset"])
|
||||
return product
|
||||
|
||||
def submit_platform(self, primary, *, count=1, ratio="4:5", platform_id="taobao"):
|
||||
product = self.platform_product()
|
||||
return enqueue_standalone_images(
|
||||
team=self.team,
|
||||
user=self.user,
|
||||
prompt="保持商品一致,生成平台电商套图",
|
||||
mode="cover",
|
||||
count=count,
|
||||
product_id=str(product.id),
|
||||
ratio=ratio,
|
||||
platform_id=platform_id,
|
||||
image_model=f"{primary.provider.name}:{primary.name}",
|
||||
dispatch=False,
|
||||
)
|
||||
|
||||
def ledger_count(self, task, ledger_type):
|
||||
return CreditLedger.objects.filter(task=task, ledger_type=ledger_type).count()
|
||||
|
||||
def test_user_selected_primary_is_first_and_success_never_queries_candidates(self):
|
||||
primary = self.model(self.provider("single-primary", 100), "user-selected")
|
||||
candidate = self.model(self.provider("single-unused", 1), "unused")
|
||||
task = self.submit(primary)[0]
|
||||
|
||||
with patch("apps.ai.routing_executor.resolve_fallback_candidates") as resolver:
|
||||
run_standalone_image_task(task_id=str(task.id))
|
||||
|
||||
resolver.assert_not_called()
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertEqual(task.model_config_id, primary.id)
|
||||
self.assertEqual(task.model_attempts.count(), 1)
|
||||
attempt = task.model_attempts.get()
|
||||
self.assertEqual(attempt.model_config_id, primary.id)
|
||||
self.assertEqual(attempt.public_model_name, "AirShelf Image")
|
||||
self.assertEqual(attempt.platform_cost, Decimal("0.5000"))
|
||||
self.provider_mocks[primary.id].image_generation.assert_called_once()
|
||||
self.assertNotIn(candidate.id, self.provider_mocks)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
|
||||
def test_primary_retry_then_dynamic_candidate_success_has_one_charge(self):
|
||||
primary = self.model(self.provider("single-fallback-primary", 100), "primary", base_cost="0.25")
|
||||
candidate = self.model(self.provider("volcano", 10), "candidate", outbound=False, base_cost="0.75")
|
||||
task = self.submit(primary, ratio="9:16")[0]
|
||||
self.provider_mocks[primary.id] = self._new_provider_mock()
|
||||
self.provider_mocks[primary.id].image_generation.side_effect = requests.ConnectionError("primary offline")
|
||||
|
||||
run_standalone_image_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
attempts = list(task.model_attempts.all())
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertEqual(task.model_config_id, primary.id)
|
||||
self.assertEqual([item.model_config_id for item in attempts], [primary.id, primary.id, candidate.id])
|
||||
self.assertEqual([item.is_retry for item in attempts], [False, True, False])
|
||||
self.assertEqual([item.is_fallback for item in attempts], [False, False, True])
|
||||
self.assertTrue(all(item.public_model_name == "AirShelf Image" for item in attempts))
|
||||
self.assertEqual(task.base_cost, Decimal("0.7500"))
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
|
||||
self.assertEqual(Asset.objects.filter(origin_task=task).count(), 1)
|
||||
|
||||
def test_direct_primary_switch_off_retries_but_does_not_call_candidate(self):
|
||||
primary = self.model(self.provider("volcano", 10), "direct-primary", outbound=False)
|
||||
candidate = self.model(self.provider("single-direct-unused", 20), "unused")
|
||||
task = self.submit(primary)[0]
|
||||
self.provider_mocks[primary.id] = self._new_provider_mock()
|
||||
self.provider_mocks[primary.id].image_generation.side_effect = requests.ConnectionError("direct offline")
|
||||
|
||||
run_standalone_image_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.FAILED)
|
||||
self.assertEqual(task.model_attempts.count(), 2)
|
||||
self.assertNotIn(candidate.id, self.provider_mocks)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
|
||||
|
||||
def test_all_candidates_fail_releases_original_reservation_once(self):
|
||||
primary = self.model(self.provider("single-all-primary", 100), "primary")
|
||||
candidate = self.model(self.provider("volcano", 10), "candidate", outbound=False)
|
||||
task = self.submit(primary)[0]
|
||||
for model in (primary, candidate):
|
||||
self.provider_mocks[model.id] = self._new_provider_mock()
|
||||
self.provider_mocks[model.id].image_generation.side_effect = requests.ConnectionError("offline")
|
||||
|
||||
run_standalone_image_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.FAILED)
|
||||
self.assertEqual(task.model_attempts.count(), 3)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
|
||||
account = CreditAccount.objects.get(team=self.team)
|
||||
self.assertEqual(account.reserved_balance, Decimal("0"))
|
||||
|
||||
def test_invalid_input_neither_retries_nor_fallbacks(self):
|
||||
primary = self.model(self.provider("single-invalid-primary", 100), "primary")
|
||||
candidate = self.model(self.provider("single-invalid-unused", 10), "unused")
|
||||
task = self.submit(primary)[0]
|
||||
self.provider_mocks[primary.id] = self._new_provider_mock()
|
||||
self.provider_mocks[primary.id].image_generation.side_effect = RuntimeError("invalid image parameter")
|
||||
|
||||
run_standalone_image_task(task_id=str(task.id))
|
||||
|
||||
self.assertEqual(task.model_attempts.count(), 1)
|
||||
self.assertEqual(task.model_attempts.get().error_type, "invalid_input")
|
||||
self.assertNotIn(candidate.id, self.provider_mocks)
|
||||
|
||||
def test_incompatible_ratio_candidate_is_excluded(self):
|
||||
primary = self.model(self.provider("single-ratio-primary", 100), "primary")
|
||||
incompatible = self.model(
|
||||
self.provider("single-ratio-incompatible", 10),
|
||||
"incompatible",
|
||||
ratios=["1:1"],
|
||||
)
|
||||
task = self.submit(primary, ratio="9:16")[0]
|
||||
self.provider_mocks[primary.id] = self._new_provider_mock()
|
||||
self.provider_mocks[primary.id].image_generation.side_effect = requests.ConnectionError("offline")
|
||||
|
||||
run_standalone_image_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.FAILED)
|
||||
self.assertEqual(task.model_attempts.count(), 2)
|
||||
self.assertNotIn(incompatible.id, self.provider_mocks)
|
||||
|
||||
def test_postprocess_failure_does_not_reinvoke_model_or_fallback(self):
|
||||
primary = self.model(self.provider("single-post-primary", 100), "primary")
|
||||
candidate = self.model(self.provider("single-post-unused", 10), "unused")
|
||||
task = self.submit(primary)[0]
|
||||
self.media_to_bytes.side_effect = RuntimeError("download failed")
|
||||
|
||||
run_standalone_image_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.FAILED)
|
||||
self.assertEqual(task.model_attempts.count(), 1)
|
||||
self.assertEqual(task.model_attempts.get().status, AIModelAttempt.Status.SUCCEEDED)
|
||||
self.provider_mocks[primary.id].image_generation.assert_called_once()
|
||||
self.assertNotIn(candidate.id, self.provider_mocks)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
|
||||
|
||||
def test_single_reference_uses_selected_edit_model_and_records_attempt(self):
|
||||
primary = self.model(
|
||||
self.provider("single-ref-primary", 20),
|
||||
"selected-edit-model",
|
||||
reference_modes=["none", "single", "multiple"],
|
||||
max_reference_images=9,
|
||||
)
|
||||
reference = self.reference_asset("single-ref")
|
||||
task = self.submit(primary, reference_image_ids=[str(reference.id)])[0]
|
||||
|
||||
run_standalone_image_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertEqual(task.model_config_id, primary.id)
|
||||
self.assertEqual(task.model_attempts.count(), 1)
|
||||
attempt = task.model_attempts.get()
|
||||
self.assertEqual(attempt.operation, "image_edit")
|
||||
self.assertEqual(attempt.request_summary["reference_images"], 1)
|
||||
call = self.provider_mocks[primary.id].image_edit.call_args
|
||||
self.assertEqual(call.kwargs["images"], ["http://example.test/single-ref.png"])
|
||||
self.assertIn("唯一依据", call.kwargs["prompt"])
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
|
||||
def test_selected_seedream_reference_model_stays_primary_and_uses_image_input(self):
|
||||
primary = self.model(
|
||||
self.provider("volcano", 10),
|
||||
"seedream-selected",
|
||||
outbound=False,
|
||||
reference_modes=["none", "single", "multiple"],
|
||||
max_reference_images=9,
|
||||
)
|
||||
reference = self.reference_asset("seedream-ref")
|
||||
self.provider_mocks[primary.id] = self._new_generation_only_provider_mock()
|
||||
|
||||
task = self.submit(primary, reference_image_ids=[str(reference.id)])[0]
|
||||
run_standalone_image_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertEqual(task.model_config_id, primary.id)
|
||||
self.assertEqual(task.model_attempts.count(), 1)
|
||||
call = self.provider_mocks[primary.id].image_generation.call_args
|
||||
self.assertEqual(call.kwargs["image"], ["http://example.test/seedream-ref.png"])
|
||||
self.assertEqual(task.model_attempts.get().public_model_name, primary.display_name)
|
||||
|
||||
def test_multiple_references_retry_then_dynamic_compatible_candidate_success(self):
|
||||
primary = self.model(
|
||||
self.provider("multi-ref-primary", 100),
|
||||
"primary-edit",
|
||||
reference_modes=["none", "single", "multiple"],
|
||||
max_reference_images=9,
|
||||
)
|
||||
candidate = self.model(
|
||||
self.provider("volcano", 10),
|
||||
"candidate-edit",
|
||||
outbound=False,
|
||||
reference_modes=["none", "single", "multiple"],
|
||||
max_reference_images=9,
|
||||
)
|
||||
references = [self.reference_asset("multi-ref-1"), self.reference_asset("multi-ref-2")]
|
||||
task = self.submit(primary, reference_image_ids=[str(item.id) for item in references])[0]
|
||||
self.provider_mocks[primary.id] = self._new_provider_mock()
|
||||
self.provider_mocks[primary.id].image_edit.side_effect = requests.ConnectionError("primary offline")
|
||||
|
||||
run_standalone_image_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
attempts = list(task.model_attempts.all())
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertEqual([item.model_config_id for item in attempts], [primary.id, primary.id, candidate.id])
|
||||
self.assertTrue(attempts[-1].is_fallback)
|
||||
self.assertTrue(all(item.operation == "image_edit" for item in attempts))
|
||||
candidate_call = self.provider_mocks[candidate.id].image_edit.call_args
|
||||
self.assertEqual(
|
||||
candidate_call.kwargs["images"],
|
||||
["http://example.test/multi-ref-1.png", "http://example.test/multi-ref-2.png"],
|
||||
)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
|
||||
|
||||
def test_multiple_references_exclude_candidate_with_insufficient_limit(self):
|
||||
primary = self.model(
|
||||
self.provider("multi-limit-primary", 100),
|
||||
"primary-edit",
|
||||
reference_modes=["none", "single", "multiple"],
|
||||
max_reference_images=9,
|
||||
)
|
||||
incompatible = self.model(
|
||||
self.provider("multi-limit-incompatible", 10),
|
||||
"single-only",
|
||||
reference_modes=["none", "single", "multiple"],
|
||||
max_reference_images=1,
|
||||
)
|
||||
references = [self.reference_asset("limit-ref-1"), self.reference_asset("limit-ref-2")]
|
||||
task = self.submit(primary, reference_image_ids=[str(item.id) for item in references])[0]
|
||||
self.provider_mocks[primary.id] = self._new_provider_mock()
|
||||
self.provider_mocks[primary.id].image_edit.side_effect = requests.ConnectionError("primary offline")
|
||||
|
||||
run_standalone_image_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.FAILED)
|
||||
self.assertEqual(task.model_attempts.count(), 2)
|
||||
self.assertNotIn(incompatible.id, self.provider_mocks)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
|
||||
|
||||
def test_tryon_selected_edit_model_is_primary_and_records_two_references(self):
|
||||
primary = self.model(
|
||||
self.provider("tryon-primary", 20),
|
||||
"selected-tryon-edit",
|
||||
reference_modes=["none", "single", "multiple"],
|
||||
max_reference_images=9,
|
||||
)
|
||||
task = self.submit_tryon(primary)[0]
|
||||
|
||||
run_standalone_image_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertEqual(task.model_config_id, primary.id)
|
||||
self.assertEqual(task.model_attempts.count(), 1)
|
||||
attempt = task.model_attempts.get()
|
||||
self.assertEqual(attempt.operation, "image_edit")
|
||||
self.assertEqual(attempt.request_summary["reference_images"], 2)
|
||||
call = self.provider_mocks[primary.id].image_edit.call_args
|
||||
self.assertEqual(
|
||||
call.kwargs["images"],
|
||||
["http://example.test/tryon-product.png", "http://example.test/tryon-model.png"],
|
||||
)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
|
||||
def test_tryon_selected_seedream_stays_primary_and_uses_image_generation(self):
|
||||
primary = self.model(
|
||||
self.provider("volcano", 10),
|
||||
"seedream-tryon-selected",
|
||||
outbound=False,
|
||||
reference_modes=["none", "single", "multiple"],
|
||||
max_reference_images=9,
|
||||
)
|
||||
self.provider_mocks[primary.id] = self._new_generation_only_provider_mock()
|
||||
task = self.submit_tryon(primary)[0]
|
||||
|
||||
run_standalone_image_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertEqual(task.model_config_id, primary.id)
|
||||
self.assertEqual(task.model_attempts.count(), 1)
|
||||
call = self.provider_mocks[primary.id].image_generation.call_args
|
||||
self.assertEqual(len(call.kwargs["image"]), 2)
|
||||
self.assertEqual(task.model_attempts.get().public_model_name, primary.display_name)
|
||||
|
||||
def test_tryon_primary_retry_then_dynamic_candidate_success_has_one_charge(self):
|
||||
primary = self.model(
|
||||
self.provider("tryon-fallback-primary", 100),
|
||||
"primary-tryon",
|
||||
reference_modes=["none", "single", "multiple"],
|
||||
max_reference_images=9,
|
||||
)
|
||||
candidate = self.model(
|
||||
self.provider("volcano", 10),
|
||||
"candidate-tryon",
|
||||
outbound=False,
|
||||
reference_modes=["none", "single", "multiple"],
|
||||
max_reference_images=9,
|
||||
)
|
||||
task = self.submit_tryon(primary)[0]
|
||||
self.provider_mocks[primary.id] = self._new_provider_mock()
|
||||
self.provider_mocks[primary.id].image_edit.side_effect = requests.ConnectionError("primary offline")
|
||||
|
||||
run_standalone_image_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
attempts = list(task.model_attempts.all())
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertEqual([item.model_config_id for item in attempts], [primary.id, primary.id, candidate.id])
|
||||
self.assertTrue(attempts[-1].is_fallback)
|
||||
self.assertTrue(all(item.request_summary["reference_images"] == 2 for item in attempts))
|
||||
self.assertTrue(all(item.public_model_name == "AirShelf Image" for item in attempts))
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
|
||||
|
||||
def test_tryon_direct_primary_failure_retries_without_fallback_and_releases(self):
|
||||
primary = self.model(
|
||||
self.provider("volcano", 10),
|
||||
"direct-tryon",
|
||||
outbound=False,
|
||||
reference_modes=["none", "single", "multiple"],
|
||||
max_reference_images=9,
|
||||
)
|
||||
candidate = self.model(
|
||||
self.provider("tryon-direct-unused", 20),
|
||||
"unused-tryon",
|
||||
reference_modes=["none", "single", "multiple"],
|
||||
max_reference_images=9,
|
||||
)
|
||||
self.provider_mocks[primary.id] = self._new_generation_only_provider_mock()
|
||||
self.provider_mocks[primary.id].image_generation.side_effect = requests.ConnectionError("direct offline")
|
||||
task = self.submit_tryon(primary)[0]
|
||||
|
||||
run_standalone_image_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.FAILED)
|
||||
self.assertEqual(task.model_attempts.count(), 2)
|
||||
self.assertNotIn(candidate.id, self.provider_mocks)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
|
||||
|
||||
def test_tryon_batch_keeps_independent_task_attempt_and_billing_chains(self):
|
||||
primary = self.model(
|
||||
self.provider("tryon-batch-primary", 20),
|
||||
"batch-tryon",
|
||||
reference_modes=["none", "single", "multiple"],
|
||||
max_reference_images=9,
|
||||
)
|
||||
tasks = self.submit_tryon(primary, count=2)
|
||||
|
||||
for task in tasks:
|
||||
run_standalone_image_task(task_id=str(task.id))
|
||||
|
||||
for task in tasks:
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertEqual(task.model_attempts.count(), 1)
|
||||
self.assertEqual(task.model_attempts.get().request_summary["reference_images"], 2)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
|
||||
|
||||
def test_platform_kit_selected_model_is_primary_and_keeps_platform_prompt(self):
|
||||
primary = self.model(
|
||||
self.provider("platform-primary", 20),
|
||||
"selected-platform-edit",
|
||||
reference_modes=["none", "single", "multiple"],
|
||||
max_reference_images=9,
|
||||
)
|
||||
task = self.submit_platform(primary, platform_id="taobao")[0]
|
||||
|
||||
run_standalone_image_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertEqual(task.model_config_id, primary.id)
|
||||
self.assertEqual(task.request_payload["platform_id"], "taobao")
|
||||
self.assertEqual(task.request_payload["platform_name"], "淘宝")
|
||||
attempt = task.model_attempts.get()
|
||||
self.assertEqual(attempt.operation, "image_edit")
|
||||
self.assertEqual(attempt.request_summary["reference_images"], 1)
|
||||
call = self.provider_mocks[primary.id].image_edit.call_args
|
||||
self.assertEqual(call.kwargs["images"], ["http://example.test/platform-product.png"])
|
||||
self.assertIn("淘宝", call.kwargs["prompt"])
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
|
||||
def test_platform_kit_selected_seedream_stays_primary(self):
|
||||
primary = self.model(
|
||||
self.provider("volcano", 10),
|
||||
"seedream-platform-selected",
|
||||
outbound=False,
|
||||
reference_modes=["none", "single", "multiple"],
|
||||
max_reference_images=9,
|
||||
)
|
||||
self.provider_mocks[primary.id] = self._new_generation_only_provider_mock()
|
||||
task = self.submit_platform(primary, platform_id="douyin")[0]
|
||||
|
||||
run_standalone_image_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertEqual(task.model_config_id, primary.id)
|
||||
self.assertEqual(task.model_attempts.count(), 1)
|
||||
call = self.provider_mocks[primary.id].image_generation.call_args
|
||||
self.assertEqual(call.kwargs["image"], ["http://example.test/platform-product.png"])
|
||||
self.assertEqual(task.model_attempts.get().public_model_name, primary.display_name)
|
||||
|
||||
def test_platform_kit_retry_then_dynamic_candidate_success_has_one_charge(self):
|
||||
primary = self.model(
|
||||
self.provider("platform-fallback-primary", 100),
|
||||
"primary-platform",
|
||||
reference_modes=["none", "single", "multiple"],
|
||||
max_reference_images=9,
|
||||
)
|
||||
candidate = self.model(
|
||||
self.provider("volcano", 10),
|
||||
"candidate-platform",
|
||||
outbound=False,
|
||||
reference_modes=["none", "single", "multiple"],
|
||||
max_reference_images=9,
|
||||
)
|
||||
task = self.submit_platform(primary)[0]
|
||||
self.provider_mocks[primary.id] = self._new_provider_mock()
|
||||
self.provider_mocks[primary.id].image_edit.side_effect = requests.ConnectionError("primary offline")
|
||||
|
||||
run_standalone_image_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
attempts = list(task.model_attempts.all())
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertEqual([item.model_config_id for item in attempts], [primary.id, primary.id, candidate.id])
|
||||
self.assertTrue(attempts[-1].is_fallback)
|
||||
self.assertTrue(all(item.request_summary["reference_images"] == 1 for item in attempts))
|
||||
self.assertTrue(all(item.public_model_name == "AirShelf Image" for item in attempts))
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
|
||||
|
||||
def test_platform_kit_batch_keeps_independent_tasks_attempts_and_billing(self):
|
||||
primary = self.model(
|
||||
self.provider("platform-batch-primary", 20),
|
||||
"batch-platform",
|
||||
reference_modes=["none", "single", "multiple"],
|
||||
max_reference_images=9,
|
||||
)
|
||||
tasks = self.submit_platform(primary, count=3, platform_id="jd")
|
||||
|
||||
for task in tasks:
|
||||
run_standalone_image_task(task_id=str(task.id))
|
||||
|
||||
self.assertEqual({task.request_payload["batch_id"] for task in tasks}, {tasks[0].request_payload["batch_id"]})
|
||||
for index, task in enumerate(tasks):
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertEqual(task.request_payload["index"], index)
|
||||
self.assertEqual(task.request_payload["platform_id"], "jd")
|
||||
self.assertEqual(task.model_attempts.count(), 1)
|
||||
self.assertEqual(task.model_attempts.get().request_summary["reference_images"], 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
|
||||
|
||||
def test_batch_without_references_keeps_legacy_path_until_its_own_step(self):
|
||||
primary = self.model(self.provider("single-batch-primary", 100), "primary")
|
||||
tasks = self.submit(primary, count=2)
|
||||
|
||||
for task in tasks:
|
||||
run_standalone_image_task(task_id=str(task.id))
|
||||
|
||||
self.assertTrue(all("model_routing_v1" not in task.request_payload for task in tasks))
|
||||
self.assertEqual(AIModelAttempt.objects.filter(task__in=tasks).count(), 0)
|
||||
self.assertEqual(self.provider_mocks[primary.id].image_generation.call_count, 2)
|
||||
@@ -0,0 +1,312 @@
|
||||
from decimal import Decimal
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import requests
|
||||
from django.test import TestCase, override_settings
|
||||
|
||||
from apps.accounts.models import Team, TeamMember, User
|
||||
from apps.ai.models import AITask, ModelConfig, ModelProvider
|
||||
from apps.ai.services import _storyboard_shot_worker, poll_storyboard
|
||||
from apps.assets.models import Asset
|
||||
from apps.billing.models import CreditAccount, CreditLedger
|
||||
from apps.products.models import Product
|
||||
from apps.projects.models import (
|
||||
Project,
|
||||
ScriptSegment,
|
||||
ScriptVersion,
|
||||
StoryboardShot,
|
||||
StoryboardShotVersion,
|
||||
)
|
||||
|
||||
|
||||
def _metadata(*, outbound=True, base_cost="0.50", max_refs=9):
|
||||
return {
|
||||
"routing": {"fallback_on_failure": outbound, "fallback_candidate": True},
|
||||
"capabilities": {
|
||||
"operations": ["image_generate", "image_edit"],
|
||||
"features": [],
|
||||
"reference_modes": ["none", "single", "multiple"],
|
||||
"max_reference_images": max_refs,
|
||||
"aspect_ratios": ["1:1", "3:4", "4:5", "9:16", "16:9"],
|
||||
},
|
||||
"pricing": {"base_cost_yuan": base_cost},
|
||||
}
|
||||
|
||||
|
||||
@override_settings(STORYBOARD_MAX_PARALLEL=4)
|
||||
class StoryboardRoutingTests(TestCase):
|
||||
def setUp(self):
|
||||
ModelConfig.objects.filter(capability=ModelConfig.Capability.IMAGE).update(
|
||||
status=ModelConfig.Status.DISABLED
|
||||
)
|
||||
self.user = User.objects.create_user(username="storyboard-routing", password="x")
|
||||
self.team = Team.objects.create(name="Storyboard Routing", owner=self.user)
|
||||
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
|
||||
CreditAccount.objects.create(team=self.team, balance=Decimal("1000"))
|
||||
product = Product.objects.create(team=self.team, created_by=self.user, title="测试商品")
|
||||
self.project = Project.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
product=product,
|
||||
name="故事板路由项目",
|
||||
)
|
||||
script = ScriptVersion.objects.create(
|
||||
project=self.project,
|
||||
title="脚本",
|
||||
content="结构化脚本",
|
||||
is_adopted=True,
|
||||
)
|
||||
self.segment = ScriptSegment.objects.create(
|
||||
script_version=script,
|
||||
sort_order=0,
|
||||
visual_prompt="人物在商品旁展示",
|
||||
)
|
||||
self.shot = StoryboardShot.objects.create(
|
||||
project=self.project,
|
||||
script_segment=self.segment,
|
||||
sort_order=0,
|
||||
status=StoryboardShot.Status.QUEUED,
|
||||
)
|
||||
self.refs = [
|
||||
{"url": "http://example.test/person.png", "label": "女主", "type": "person"},
|
||||
{"url": "http://example.test/product.png", "label": "商品", "type": "product"},
|
||||
]
|
||||
self.provider_mocks = {}
|
||||
patch("apps.ai.services.get_image_provider", side_effect=self._provider_for).start()
|
||||
patch(
|
||||
"apps.ai.services._storyboard_reference_images",
|
||||
side_effect=lambda project, segment: list(self.refs),
|
||||
).start()
|
||||
patch(
|
||||
"apps.ai.services.build_storyboard_frame_prompt_refs",
|
||||
return_value="带参考图编号与锁定约束的故事板提示词",
|
||||
).start()
|
||||
patch(
|
||||
"apps.ai.services.build_storyboard_frame_prompt",
|
||||
return_value="无参考图故事板提示词",
|
||||
).start()
|
||||
patch("apps.ai.services._store_generated_media", side_effect=self._store_media).start()
|
||||
patch("apps.assets.review.submit_asset_for_review").start()
|
||||
patch("apps.ai.services.notify_generation_failure").start()
|
||||
self.addCleanup(patch.stopall)
|
||||
|
||||
def provider(self, name, priority):
|
||||
return ModelProvider.objects.create(
|
||||
name=name,
|
||||
display_name=name,
|
||||
status=ModelProvider.Status.ACTIVE,
|
||||
metadata={"routing": {"fallback_priority": priority}},
|
||||
)
|
||||
|
||||
def model(self, provider, name, *, outbound=True, base_cost="0.50", max_refs=9):
|
||||
return ModelConfig.objects.create(
|
||||
provider=provider,
|
||||
name=name,
|
||||
display_name=name,
|
||||
capability=ModelConfig.Capability.IMAGE,
|
||||
endpoint="images/generations",
|
||||
unit_price=Decimal("20"),
|
||||
status=ModelConfig.Status.ACTIVE,
|
||||
metadata=_metadata(outbound=outbound, base_cost=base_cost, max_refs=max_refs),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _new_provider_mock():
|
||||
provider = Mock()
|
||||
response = {"data": [{"url": "http://example.test/storyboard.png"}]}
|
||||
provider.image_edit.return_value = response
|
||||
provider.image_generation.return_value = response
|
||||
provider.extract_first_media_url.side_effect = lambda value: value["data"][0]["url"]
|
||||
return provider
|
||||
|
||||
@staticmethod
|
||||
def _new_generation_only_provider_mock():
|
||||
provider = Mock(spec=["image_generation", "extract_first_media_url"])
|
||||
provider.image_generation.return_value = {
|
||||
"data": [{"url": "http://example.test/storyboard.png"}]
|
||||
}
|
||||
provider.extract_first_media_url.side_effect = lambda value: value["data"][0]["url"]
|
||||
return provider
|
||||
|
||||
def _provider_for(self, model):
|
||||
return self.provider_mocks.setdefault(model.id, self._new_provider_mock())
|
||||
|
||||
def _store_media(self, **kwargs):
|
||||
return Asset.objects.create(
|
||||
team=kwargs["team"],
|
||||
created_by=kwargs["user"],
|
||||
name=kwargs["name"],
|
||||
asset_type=kwargs["asset_type"],
|
||||
source=Asset.Source.AI_GENERATED,
|
||||
category=kwargs["category"],
|
||||
origin_task=kwargs["task"],
|
||||
)
|
||||
|
||||
def enqueue(self, *, run=True):
|
||||
with patch("threading.Thread"):
|
||||
result = poll_storyboard(project=self.project, user=self.user)
|
||||
self.assertEqual(result["status"], "generating")
|
||||
task = AITask.objects.filter(project=self.project, task_type=AITask.Type.STORYBOARD).latest(
|
||||
"created_at"
|
||||
)
|
||||
if run:
|
||||
_storyboard_shot_worker(str(task.id), str(self.shot.id), str(self.user.id))
|
||||
return task
|
||||
|
||||
@staticmethod
|
||||
def ledger_count(task, ledger_type):
|
||||
return CreditLedger.objects.filter(task=task, ledger_type=ledger_type).count()
|
||||
|
||||
def test_multi_reference_success_records_9_16_and_adopts_one_shot_version(self):
|
||||
primary = self.model(self.provider("storyboard-primary", 20), "storyboard-primary")
|
||||
task = self.enqueue()
|
||||
|
||||
task.refresh_from_db()
|
||||
self.shot.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertTrue(task.request_payload["model_routing_v1"])
|
||||
self.assertEqual(task.base_cost, Decimal("0.5000"))
|
||||
attempt = task.model_attempts.get()
|
||||
self.assertEqual(attempt.model_config_id, primary.id)
|
||||
self.assertEqual(attempt.operation, "image_edit")
|
||||
self.assertEqual(attempt.request_summary["reference_images"], 2)
|
||||
self.assertEqual(attempt.request_summary["aspect_ratio"], "9:16")
|
||||
self.assertEqual(attempt.request_summary["storyboard_shot"], str(self.shot.id))
|
||||
call = self.provider_mocks[primary.id].image_edit.call_args
|
||||
self.assertEqual(
|
||||
call.kwargs["images"],
|
||||
["http://example.test/person.png", "http://example.test/product.png"],
|
||||
)
|
||||
self.assertEqual(call.kwargs["prompt"], "带参考图编号与锁定约束的故事板提示词")
|
||||
self.assertEqual(call.kwargs["size"], "1024x1536")
|
||||
self.assertIsNotNone(self.shot.adopted_version_id)
|
||||
self.assertEqual(self.shot.adopted_version.task_id, task.id)
|
||||
self.assertEqual(self.shot.versions.count(), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
|
||||
def test_retry_then_dynamic_candidate_success_charges_once(self):
|
||||
primary = self.model(
|
||||
self.provider("storyboard-fallback-primary", 100),
|
||||
"storyboard-primary",
|
||||
base_cost="0.25",
|
||||
)
|
||||
candidate = self.model(
|
||||
self.provider("volcano", 10),
|
||||
"storyboard-candidate",
|
||||
outbound=False,
|
||||
base_cost="0.75",
|
||||
)
|
||||
self.provider_mocks[primary.id] = self._new_provider_mock()
|
||||
self.provider_mocks[primary.id].image_edit.side_effect = requests.ConnectionError("offline")
|
||||
task = self.enqueue()
|
||||
|
||||
task.refresh_from_db()
|
||||
attempts = list(task.model_attempts.all())
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertEqual([a.model_config_id for a in attempts], [primary.id, primary.id, candidate.id])
|
||||
self.assertTrue(attempts[-1].is_fallback)
|
||||
self.assertEqual(task.base_cost, Decimal("0.7500"))
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
|
||||
|
||||
def test_all_candidates_fail_releases_and_keeps_old_adopted_version(self):
|
||||
primary = self.model(self.provider("storyboard-all-primary", 100), "storyboard-primary")
|
||||
candidate = self.model(
|
||||
self.provider("volcano", 10), "storyboard-candidate", outbound=False
|
||||
)
|
||||
old_asset = Asset.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
name="旧故事板",
|
||||
asset_type=Asset.Type.IMAGE,
|
||||
source=Asset.Source.AI_GENERATED,
|
||||
category=Asset.Category.STORYBOARD,
|
||||
)
|
||||
old_version = StoryboardShotVersion.objects.create(
|
||||
shot=self.shot,
|
||||
asset=old_asset,
|
||||
prompt="旧画面",
|
||||
is_adopted=True,
|
||||
)
|
||||
self.shot.adopted_version = old_version
|
||||
self.shot.save(update_fields=["adopted_version", "updated_at"])
|
||||
for model in (primary, candidate):
|
||||
self.provider_mocks[model.id] = self._new_provider_mock()
|
||||
self.provider_mocks[model.id].image_edit.side_effect = requests.ConnectionError("offline")
|
||||
|
||||
task = self.enqueue()
|
||||
|
||||
task.refresh_from_db()
|
||||
self.shot.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.FAILED)
|
||||
self.assertEqual(task.model_attempts.count(), 3)
|
||||
self.assertEqual(self.shot.adopted_version_id, old_version.id)
|
||||
self.assertEqual(self.shot.versions.count(), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
|
||||
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
|
||||
|
||||
def test_no_reference_uses_image_generation_and_keeps_vertical_capability(self):
|
||||
primary = self.model(self.provider("storyboard-no-ref", 20), "storyboard-no-ref")
|
||||
self.refs = []
|
||||
task = self.enqueue()
|
||||
|
||||
task.refresh_from_db()
|
||||
attempt = task.model_attempts.get()
|
||||
self.assertEqual(attempt.operation, "image_generate")
|
||||
self.assertEqual(attempt.request_summary["reference_images"], 0)
|
||||
self.assertEqual(attempt.request_summary["aspect_ratio"], "9:16")
|
||||
call = self.provider_mocks[primary.id].image_generation.call_args
|
||||
self.assertEqual(call.kwargs["prompt"], "无参考图故事板提示词")
|
||||
self.assertNotIn("size", call.kwargs)
|
||||
|
||||
def test_direct_primary_uses_image_generation_with_all_references(self):
|
||||
primary = self.model(
|
||||
self.provider("volcano", 10), "seedream-storyboard", outbound=False
|
||||
)
|
||||
self.provider_mocks[primary.id] = self._new_generation_only_provider_mock()
|
||||
task = self.enqueue()
|
||||
|
||||
task.refresh_from_db()
|
||||
call = self.provider_mocks[primary.id].image_generation.call_args
|
||||
self.assertEqual(
|
||||
call.kwargs["image"],
|
||||
["http://example.test/person.png", "http://example.test/product.png"],
|
||||
)
|
||||
self.assertEqual(call.kwargs["size"], "1024x1536")
|
||||
self.assertEqual(task.model_attempts.get().public_model_name, primary.display_name)
|
||||
|
||||
def test_poll_creates_one_independent_task_and_reservation_per_shot(self):
|
||||
self.model(self.provider("storyboard-batch", 20), "storyboard-batch")
|
||||
second_segment = ScriptSegment.objects.create(
|
||||
script_version=self.segment.script_version,
|
||||
sort_order=1,
|
||||
visual_prompt="第二镜",
|
||||
)
|
||||
second_shot = StoryboardShot.objects.create(
|
||||
project=self.project,
|
||||
script_segment=second_segment,
|
||||
sort_order=1,
|
||||
status=StoryboardShot.Status.QUEUED,
|
||||
)
|
||||
|
||||
with patch("threading.Thread"):
|
||||
result = poll_storyboard(project=self.project, user=self.user)
|
||||
|
||||
self.assertEqual(result, {"status": "generating", "done": 0, "total": 2})
|
||||
tasks = list(
|
||||
AITask.objects.filter(project=self.project, task_type=AITask.Type.STORYBOARD).order_by(
|
||||
"created_at"
|
||||
)
|
||||
)
|
||||
self.assertEqual(len(tasks), 2)
|
||||
self.assertEqual(
|
||||
{task.request_payload["storyboard_shot"] for task in tasks},
|
||||
{str(self.shot.id), str(second_shot.id)},
|
||||
)
|
||||
self.assertTrue(all(task.request_payload["model_routing_v1"] for task in tasks))
|
||||
self.assertTrue(all(task.base_cost == Decimal("0") for task in tasks))
|
||||
self.assertTrue(all(self.ledger_count(task, CreditLedger.Type.RESERVE) == 1 for task in tasks))
|
||||
@@ -0,0 +1,278 @@
|
||||
from copy import deepcopy
|
||||
from datetime import timedelta
|
||||
from decimal import Decimal
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import requests
|
||||
from django.conf import settings
|
||||
from django.test import TestCase, override_settings
|
||||
from django.utils import timezone
|
||||
|
||||
from apps.accounts.models import Team, TeamMember, User
|
||||
from apps.ai.models import AITask, ModelConfig, ModelProvider
|
||||
from apps.ai.services import poll_video_segment, submit_video_segment
|
||||
from apps.assets.models import Asset
|
||||
from apps.billing.models import CreditAccount, CreditLedger
|
||||
from apps.products.models import Product
|
||||
from apps.projects.models import Project, VideoSegment
|
||||
|
||||
|
||||
def _fast_routing_policy():
|
||||
policy = deepcopy(settings.MODEL_ROUTING_POLICY)
|
||||
policy["jitter_ratio"] = 0
|
||||
policy["video"]["submit_retry_delays"] = [0]
|
||||
return policy
|
||||
|
||||
|
||||
def _video_metadata(*, outbound, candidate=True, price=46):
|
||||
return {
|
||||
"routing": {
|
||||
"fallback_on_failure": outbound,
|
||||
"fallback_candidate": candidate,
|
||||
},
|
||||
"capabilities": {
|
||||
"operations": ["video_generate"],
|
||||
"features": ["generate_audio"],
|
||||
"max_reference_images": 9,
|
||||
"max_reference_videos": 3,
|
||||
"max_reference_audios": 3,
|
||||
"aspect_ratios": ["9:16"],
|
||||
"resolutions": ["720p"],
|
||||
"durations": [15],
|
||||
},
|
||||
"pricing": {
|
||||
"unit": "cny_per_million_tokens",
|
||||
"default": {"no_ref_video": price, "with_ref_video": price},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@override_settings(MODEL_ROUTING_POLICY=_fast_routing_policy())
|
||||
class VideoSegmentRoutingTests(TestCase):
|
||||
def setUp(self):
|
||||
ModelConfig.objects.filter(capability=ModelConfig.Capability.VIDEO).update(
|
||||
status=ModelConfig.Status.DISABLED,
|
||||
is_default=False,
|
||||
)
|
||||
self.user = User.objects.create_user(username="video-routing", password="x")
|
||||
self.team = Team.objects.create(name="Video Routing", owner=self.user)
|
||||
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
|
||||
CreditAccount.objects.create(team=self.team, balance=Decimal("10000"))
|
||||
product = Product.objects.create(team=self.team, created_by=self.user, title="Product")
|
||||
self.project = Project.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
product=product,
|
||||
name="Video routing project",
|
||||
)
|
||||
self.segment = VideoSegment.objects.create(
|
||||
project=self.project,
|
||||
sort_order=0,
|
||||
target_duration_seconds=15,
|
||||
)
|
||||
self.provider_mocks = {}
|
||||
patch("apps.ai.services.get_video_provider", side_effect=self._provider_for).start()
|
||||
patch("apps.ai.services._video_reference_images", return_value=[]).start()
|
||||
patch("apps.ai.services.build_video_segment_prompt", return_value="video prompt").start()
|
||||
patch("apps.ai.services.notify_generation_failure").start()
|
||||
self.addCleanup(patch.stopall)
|
||||
|
||||
def provider(self, name, priority):
|
||||
return ModelProvider.objects.create(
|
||||
name=name,
|
||||
display_name=name,
|
||||
status=ModelProvider.Status.ACTIVE,
|
||||
base_url="https://video.example/v1",
|
||||
metadata={"routing": {"fallback_priority": priority}},
|
||||
)
|
||||
|
||||
def model(self, provider, name, *, outbound, price=46, default=False):
|
||||
return ModelConfig.objects.create(
|
||||
provider=provider,
|
||||
name=name,
|
||||
display_name=name,
|
||||
capability=ModelConfig.Capability.VIDEO,
|
||||
endpoint="contents/generations/tasks",
|
||||
unit_price=Decimal("1"),
|
||||
status=ModelConfig.Status.ACTIVE,
|
||||
is_default=default,
|
||||
metadata=_video_metadata(outbound=outbound, price=price),
|
||||
)
|
||||
|
||||
def _provider_for(self, model):
|
||||
return self.provider_mocks.setdefault(model.id, Mock())
|
||||
|
||||
def _submit(self):
|
||||
submit_video_segment(video_segment=self.segment, user=self.user, prompt="test")
|
||||
return AITask.objects.get(task_type=AITask.Type.VIDEO_SEGMENT, team=self.team)
|
||||
|
||||
@staticmethod
|
||||
def ledger_count(task, ledger_type):
|
||||
return CreditLedger.objects.filter(task=task, ledger_type=ledger_type).count()
|
||||
|
||||
@staticmethod
|
||||
def _asset(task):
|
||||
return Asset.objects.create(
|
||||
team=task.team,
|
||||
created_by=task.created_by,
|
||||
name="video clip",
|
||||
asset_type=Asset.Type.VIDEO,
|
||||
source=Asset.Source.AI_GENERATED,
|
||||
category=Asset.Category.VIDEO_CLIP,
|
||||
origin_task=task,
|
||||
)
|
||||
|
||||
def test_first_submit_success_creates_one_attempt_and_one_reservation(self):
|
||||
primary = self.model(self.provider("video-primary-ok", 20), "video-primary-ok", outbound=True, default=True)
|
||||
provider = self._provider_for(primary)
|
||||
provider.create_video_task.return_value = {"id": "remote-primary", "status": "queued"}
|
||||
|
||||
task = self._submit()
|
||||
|
||||
task.refresh_from_db()
|
||||
self.segment.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.SUBMITTED)
|
||||
self.assertEqual(task.provider_task_id, "remote-primary")
|
||||
self.assertEqual(task.request_payload["actual_model_config_id"], str(primary.id))
|
||||
self.assertEqual(task.model_attempts.count(), 1)
|
||||
self.assertEqual(task.model_attempts.get().provider_task_id, "remote-primary")
|
||||
self.assertEqual(self.segment.status, VideoSegment.Status.RUNNING)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
|
||||
|
||||
@patch("apps.ai.services._store_generated_media")
|
||||
def test_retry_then_dynamic_fallback_polls_actual_model_and_charges_once(self, store_media):
|
||||
primary = self.model(self.provider("video-primary-fail", 20), "video-primary-fail", outbound=True, price=46, default=True)
|
||||
candidate = self.model(self.provider("video-candidate-ok", 30), "video-candidate-ok", outbound=False, price=23)
|
||||
primary_provider = self._provider_for(primary)
|
||||
fallback_provider = self._provider_for(candidate)
|
||||
primary_provider.create_video_task.side_effect = requests.ConnectionError("primary offline")
|
||||
fallback_provider.create_video_task.return_value = {"id": "remote-fallback", "status": "queued"}
|
||||
|
||||
task = self._submit()
|
||||
store_media.return_value = self._asset(task)
|
||||
fallback_provider.poll_video_task.return_value = {
|
||||
"status": "succeeded",
|
||||
"usage": {"total_tokens": 300000},
|
||||
"content": {"video_url": "https://video.example/result.mp4"},
|
||||
}
|
||||
fallback_provider.extract_first_media_url.return_value = "https://video.example/result.mp4"
|
||||
|
||||
version = poll_video_segment(video_segment=self.segment, user=self.user)
|
||||
|
||||
task.refresh_from_db()
|
||||
attempts = list(task.model_attempts.order_by("sequence"))
|
||||
self.assertIsNotNone(version)
|
||||
self.assertEqual(
|
||||
[attempt.model_config_id for attempt in attempts],
|
||||
[primary.id, primary.id, candidate.id],
|
||||
)
|
||||
self.assertTrue(attempts[-1].is_fallback)
|
||||
self.assertEqual(task.model_config_id, primary.id)
|
||||
self.assertEqual(task.request_payload["actual_model_config_id"], str(candidate.id))
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertFalse(primary_provider.poll_video_task.called)
|
||||
fallback_provider.poll_video_task.assert_called_once_with(
|
||||
endpoint=candidate.endpoint,
|
||||
provider_task_id="remote-fallback",
|
||||
timeout=60.0,
|
||||
)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
|
||||
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
|
||||
|
||||
def test_direct_model_with_outbound_disabled_never_uses_candidate(self):
|
||||
primary = self.model(self.provider("ark", 10), "video-direct-locked", outbound=False, default=True)
|
||||
candidate = self.model(self.provider("video-unused", 20), "video-unused", outbound=False)
|
||||
primary_provider = self._provider_for(primary)
|
||||
unused_provider = self._provider_for(candidate)
|
||||
primary_provider.create_video_task.side_effect = requests.ConnectionError("direct offline")
|
||||
unused_provider.create_video_task.return_value = {"id": "must-not-run"}
|
||||
|
||||
with self.assertRaises(requests.ConnectionError):
|
||||
self._submit()
|
||||
|
||||
task = AITask.objects.get(task_type=AITask.Type.VIDEO_SEGMENT, team=self.team)
|
||||
self.assertEqual(
|
||||
list(task.model_attempts.values_list("model_config_id", flat=True)),
|
||||
[primary.id, primary.id],
|
||||
)
|
||||
self.assertFalse(unused_provider.create_video_task.called)
|
||||
self.assertEqual(task.status, AITask.Status.FAILED)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
|
||||
|
||||
def test_submit_read_timeout_is_state_unknown_and_is_never_retried(self):
|
||||
primary = self.model(self.provider("video-timeout", 20), "video-timeout", outbound=True, default=True)
|
||||
candidate = self.model(self.provider("video-timeout-unused", 30), "video-timeout-unused", outbound=False)
|
||||
primary_provider = self._provider_for(primary)
|
||||
unused_provider = self._provider_for(candidate)
|
||||
primary_provider.create_video_task.side_effect = requests.ReadTimeout("response lost")
|
||||
|
||||
with self.assertRaisesRegex(RuntimeError, "state unknown|状态未知"):
|
||||
self._submit()
|
||||
|
||||
task = AITask.objects.get(task_type=AITask.Type.VIDEO_SEGMENT, team=self.team)
|
||||
self.assertEqual(task.model_attempts.count(), 1)
|
||||
self.assertFalse(unused_provider.create_video_task.called)
|
||||
self.assertEqual(task.status, AITask.Status.FAILED)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
|
||||
|
||||
def test_submit_response_without_remote_id_is_never_retried(self):
|
||||
primary = self.model(self.provider("video-no-id", 20), "video-no-id", outbound=True, default=True)
|
||||
candidate = self.model(self.provider("video-no-id-unused", 30), "video-no-id-unused", outbound=False)
|
||||
primary_provider = self._provider_for(primary)
|
||||
unused_provider = self._provider_for(candidate)
|
||||
primary_provider.create_video_task.return_value = {"status": "queued"}
|
||||
|
||||
with self.assertRaisesRegex(RuntimeError, "state unknown|状态未知"):
|
||||
self._submit()
|
||||
|
||||
task = AITask.objects.get(task_type=AITask.Type.VIDEO_SEGMENT, team=self.team)
|
||||
attempt = task.model_attempts.get()
|
||||
self.assertTrue(attempt.response_summary["provider_task_id_missing"])
|
||||
self.assertFalse(unused_provider.create_video_task.called)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
|
||||
|
||||
def test_poll_network_error_keeps_task_and_reservation_for_later_poll(self):
|
||||
primary = self.model(self.provider("video-poll-unknown", 20), "video-poll-unknown", outbound=True, default=True)
|
||||
provider = self._provider_for(primary)
|
||||
provider.create_video_task.return_value = {"id": "remote-poll", "status": "queued"}
|
||||
task = self._submit()
|
||||
provider.poll_video_task.side_effect = requests.ConnectionError("poll offline")
|
||||
|
||||
with self.assertRaises(requests.ConnectionError):
|
||||
poll_video_segment(video_segment=self.segment, user=self.user)
|
||||
|
||||
task.refresh_from_db()
|
||||
self.segment.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.SUBMITTED)
|
||||
self.assertEqual(self.segment.status, VideoSegment.Status.RUNNING)
|
||||
self.assertEqual(task.model_attempts.count(), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
|
||||
self.assertGreater(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
|
||||
|
||||
def test_generation_timeout_releases_without_polling_or_resubmitting(self):
|
||||
primary = self.model(self.provider("video-generation-timeout", 20), "video-generation-timeout", outbound=True, default=True)
|
||||
provider = self._provider_for(primary)
|
||||
provider.create_video_task.return_value = {"id": "remote-slow", "status": "queued"}
|
||||
task = self._submit()
|
||||
task.submitted_at = timezone.now() - timedelta(seconds=1900)
|
||||
task.save(update_fields=["submitted_at", "updated_at"])
|
||||
|
||||
result = poll_video_segment(video_segment=self.segment, user=self.user)
|
||||
|
||||
task.refresh_from_db()
|
||||
self.segment.refresh_from_db()
|
||||
self.assertIsNone(result)
|
||||
self.assertEqual(task.status, AITask.Status.FAILED)
|
||||
self.assertEqual(self.segment.status, VideoSegment.Status.FAILED)
|
||||
self.assertFalse(provider.poll_video_task.called)
|
||||
self.assertEqual(provider.create_video_task.call_count, 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
|
||||
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
|
||||
@@ -0,0 +1,241 @@
|
||||
from decimal import Decimal
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import requests
|
||||
from django.test import SimpleTestCase, TestCase
|
||||
|
||||
from apps.accounts.models import Team, TeamMember, User
|
||||
from apps.ai.models import AITask, ModelConfig, ModelProvider
|
||||
from apps.ai.providers.openai_compatible import OpenAICompatibleProvider
|
||||
from apps.ai.services import DEFAULT_VOICEOVER_VOICE, synthesize_project_voiceover
|
||||
from apps.assets.models import Asset
|
||||
from apps.assets.storage import StoredObject
|
||||
from apps.billing.models import CreditAccount, CreditLedger
|
||||
from apps.products.models import Product
|
||||
from apps.projects.models import Project
|
||||
|
||||
|
||||
def _metadata(*, outbound, candidate=True, voice="BV700_streaming", base_cost="0.10"):
|
||||
return {
|
||||
"routing": {
|
||||
"fallback_on_failure": outbound,
|
||||
"fallback_candidate": candidate,
|
||||
},
|
||||
"capabilities": {
|
||||
"operations": ["tts"],
|
||||
"features": [],
|
||||
"languages": ["zh-CN"],
|
||||
"voice_map": {DEFAULT_VOICEOVER_VOICE: voice},
|
||||
"max_chars": 10000,
|
||||
"speed_range": [0.5, 2.0],
|
||||
"output_formats": ["mp3"],
|
||||
},
|
||||
"pricing": {
|
||||
"chars_per_unit": 500,
|
||||
"points_per_unit": 10,
|
||||
"min_units": 1,
|
||||
"base_cost_yuan_per_unit": base_cost,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
class VoiceoverRoutingTests(TestCase):
|
||||
def setUp(self):
|
||||
ModelConfig.objects.filter(capability=ModelConfig.Capability.AUDIO).update(
|
||||
status=ModelConfig.Status.DISABLED
|
||||
)
|
||||
self.user = User.objects.create_user(username="voice-routing", password="x")
|
||||
self.team = Team.objects.create(name="Voice Routing", owner=self.user)
|
||||
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
|
||||
CreditAccount.objects.create(team=self.team, balance=Decimal("1000"))
|
||||
product = Product.objects.create(team=self.team, created_by=self.user, title="测试商品")
|
||||
self.project = Project.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
product=product,
|
||||
name="配音路由项目",
|
||||
)
|
||||
self.direct_provider, _ = ModelProvider.objects.get_or_create(
|
||||
name="ark",
|
||||
defaults={"display_name": "ARK"},
|
||||
)
|
||||
self.direct_provider.status = ModelProvider.Status.ACTIVE
|
||||
self.direct_provider.metadata = {"routing": {"fallback_priority": 10}}
|
||||
self.direct_provider.save(update_fields=["status", "metadata", "updated_at"])
|
||||
self.provider_mocks = {}
|
||||
self.direct_mock = Mock()
|
||||
self.direct_mock.configured = True
|
||||
self.direct_mock.synthesize.return_value = (b"direct-mp3", 1200)
|
||||
patch("apps.ai.services.get_audio_provider", side_effect=self._provider_for).start()
|
||||
stored = StoredObject(
|
||||
object_key="qa-voice.mp3",
|
||||
bucket="qa",
|
||||
content_type="audio/mpeg",
|
||||
size_bytes=10,
|
||||
)
|
||||
patch("apps.ai.services.TosStorage.upload_fileobj", return_value=stored).start()
|
||||
self.addCleanup(patch.stopall)
|
||||
|
||||
def provider(self, name, priority):
|
||||
return ModelProvider.objects.create(
|
||||
name=name,
|
||||
display_name=name,
|
||||
status=ModelProvider.Status.ACTIVE,
|
||||
base_url="https://voice.example/v1",
|
||||
metadata={"routing": {"fallback_priority": priority}},
|
||||
)
|
||||
|
||||
def model(self, provider, name, *, outbound, voice="BV700_streaming", base_cost="0.10"):
|
||||
return ModelConfig.objects.create(
|
||||
provider=provider,
|
||||
name=name,
|
||||
display_name=name,
|
||||
capability=ModelConfig.Capability.AUDIO,
|
||||
endpoint="audio/speech",
|
||||
unit_price=Decimal("10"),
|
||||
status=ModelConfig.Status.ACTIVE,
|
||||
metadata=_metadata(outbound=outbound, voice=voice, base_cost=base_cost),
|
||||
)
|
||||
|
||||
def _provider_for(self, model):
|
||||
if model.provider_id == self.direct_provider.id:
|
||||
return self.direct_mock
|
||||
return self.provider_mocks.setdefault(model.id, Mock())
|
||||
|
||||
@staticmethod
|
||||
def ledger_count(task, ledger_type):
|
||||
return CreditLedger.objects.filter(task=task, ledger_type=ledger_type).count()
|
||||
|
||||
def synthesize(self, items):
|
||||
return synthesize_project_voiceover(
|
||||
project=self.project,
|
||||
user=self.user,
|
||||
items=items,
|
||||
voice_type=DEFAULT_VOICEOVER_VOICE,
|
||||
speed_ratio=1.0,
|
||||
)
|
||||
|
||||
def test_two_cues_create_two_attempts_but_one_reservation_and_charge(self):
|
||||
primary = self.model(self.direct_provider, "voice-direct-success", outbound=False)
|
||||
|
||||
voiceover = self.synthesize(
|
||||
[{"index": 0, "text": "第一镜旁白"}, {"index": 1, "text": "第二镜旁白"}]
|
||||
)
|
||||
|
||||
task = AITask.objects.get(task_type=AITask.Type.VOICEOVER)
|
||||
attempts = list(task.model_attempts.all())
|
||||
self.assertEqual(len(voiceover["items"]), 2)
|
||||
self.assertEqual([attempt.model_config_id for attempt in attempts], [primary.id, primary.id])
|
||||
self.assertTrue(all(attempt.operation == "tts" for attempt in attempts))
|
||||
self.assertEqual(self.direct_mock.synthesize.call_count, 2)
|
||||
self.assertEqual(Asset.objects.filter(origin_task=task).count(), 2)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
|
||||
|
||||
def test_direct_model_with_outbound_disabled_retries_only_primary_and_releases(self):
|
||||
primary = self.model(self.direct_provider, "voice-direct-locked", outbound=False)
|
||||
candidate = self.model(
|
||||
self.provider("voice-unused-candidate", 20),
|
||||
"voice-unused",
|
||||
outbound=False,
|
||||
voice="alloy",
|
||||
)
|
||||
self.direct_mock.synthesize.side_effect = requests.ConnectionError("offline")
|
||||
unused = Mock()
|
||||
unused.synthesize.return_value = (b"unused", 1000)
|
||||
self.provider_mocks[candidate.id] = unused
|
||||
|
||||
with self.assertRaises(requests.ConnectionError):
|
||||
self.synthesize([{"index": 0, "text": "旁白"}])
|
||||
|
||||
task = AITask.objects.get(task_type=AITask.Type.VOICEOVER)
|
||||
attempts = list(task.model_attempts.all())
|
||||
self.assertEqual([attempt.model_config_id for attempt in attempts], [primary.id, primary.id])
|
||||
self.assertFalse(unused.synthesize.called)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
|
||||
|
||||
def test_openai_candidate_enters_dynamically_and_uses_its_voice_map(self):
|
||||
primary = self.model(self.direct_provider, "voice-dynamic-primary", outbound=True)
|
||||
candidate = self.model(
|
||||
self.provider("voice-dynamic-candidate", 20),
|
||||
"voice-dynamic-fallback",
|
||||
outbound=False,
|
||||
voice="alloy",
|
||||
base_cost="0.25",
|
||||
)
|
||||
self.direct_mock.synthesize.side_effect = requests.ConnectionError("offline")
|
||||
fallback = Mock()
|
||||
fallback.synthesize.return_value = (b"fallback-mp3", 1500)
|
||||
self.provider_mocks[candidate.id] = fallback
|
||||
|
||||
voiceover = self.synthesize([{"index": 0, "text": "候选配音"}])
|
||||
|
||||
task = AITask.objects.get(task_type=AITask.Type.VOICEOVER)
|
||||
attempts = list(task.model_attempts.all())
|
||||
self.assertEqual(len(voiceover["items"]), 1)
|
||||
self.assertEqual(
|
||||
[attempt.model_config_id for attempt in attempts],
|
||||
[primary.id, primary.id, candidate.id],
|
||||
)
|
||||
self.assertTrue(attempts[-1].is_fallback)
|
||||
self.assertEqual(fallback.synthesize.call_args.kwargs["voice_type"], "alloy")
|
||||
self.assertEqual(fallback.synthesize.call_args.kwargs["model"], candidate.name)
|
||||
self.assertEqual(fallback.synthesize.call_args.kwargs["endpoint"], "audio/speech")
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0)
|
||||
|
||||
def test_all_candidates_fail_release_once_and_create_no_audio_assets(self):
|
||||
primary = self.model(self.direct_provider, "voice-all-primary", outbound=True)
|
||||
candidate = self.model(
|
||||
self.provider("voice-all-candidate", 20),
|
||||
"voice-all-fallback",
|
||||
outbound=False,
|
||||
voice="alloy",
|
||||
)
|
||||
self.direct_mock.synthesize.side_effect = requests.ConnectionError("offline")
|
||||
failed = Mock()
|
||||
failed.synthesize.side_effect = requests.ConnectionError("also offline")
|
||||
self.provider_mocks[candidate.id] = failed
|
||||
|
||||
with self.assertRaises(requests.ConnectionError):
|
||||
self.synthesize([{"index": 0, "text": "失败旁白"}])
|
||||
|
||||
task = AITask.objects.get(task_type=AITask.Type.VOICEOVER)
|
||||
self.assertEqual(task.status, AITask.Status.FAILED)
|
||||
self.assertEqual(Asset.objects.filter(origin_task=task).count(), 0)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
|
||||
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
|
||||
|
||||
|
||||
class OpenAICompatibleSpeechTests(SimpleTestCase):
|
||||
@patch("apps.ai.providers.openai_compatible.requests.post")
|
||||
def test_standard_audio_speech_request(self, post):
|
||||
response = Mock()
|
||||
response.ok = True
|
||||
response.content = b"mp3-bytes"
|
||||
response.headers = {"x-audio-duration-ms": "1234"}
|
||||
post.return_value = response
|
||||
provider = OpenAICompatibleProvider(base_url="https://voice.example/v1", api_key="secret")
|
||||
|
||||
audio, duration_ms = provider.synthesize(
|
||||
model="tts-model",
|
||||
text="你好",
|
||||
voice_type="alloy",
|
||||
speed_ratio=1.2,
|
||||
endpoint="audio/speech",
|
||||
timeout=9,
|
||||
)
|
||||
|
||||
self.assertEqual(audio, b"mp3-bytes")
|
||||
self.assertEqual(duration_ms, 1234)
|
||||
call = post.call_args
|
||||
self.assertEqual(call.args[0], "https://voice.example/v1/audio/speech")
|
||||
self.assertEqual(call.kwargs["json"]["voice"], "alloy")
|
||||
self.assertEqual(call.kwargs["json"]["response_format"], "mp3")
|
||||
self.assertEqual(call.kwargs["timeout"], 9)
|
||||
@@ -902,9 +902,9 @@ class StandaloneImageReferenceTests(TestCase):
|
||||
self.assertIn("用户要求是最高优先级", used_prompt)
|
||||
prov.image_generation.assert_not_called()
|
||||
|
||||
def test_refs_force_image_edit_model_when_volcano_selected(self):
|
||||
"""用户选了火山 Seedream(只有图生图、不会真锁主体)却传了参考图时,必须自动切到支持
|
||||
image_edit 的 gpt-image-2——否则参考图形同虚设。回归保护「选火山+传参考图=完全不参考」。"""
|
||||
def test_refs_keep_user_selected_volcano_model_as_primary(self):
|
||||
"""图片创作上传参考图后仍必须尊重用户选择的 Seedream;它可走带 image 的图生图,
|
||||
只有真实调用失败后才允许统一路由按配置 Fallback,不能在建任务阶段静默换成 gpt-image。"""
|
||||
from apps.ai.models import ModelProvider
|
||||
|
||||
vp, _ = ModelProvider.objects.get_or_create(name="volcengine", defaults={"display_name": "火山"})
|
||||
@@ -919,8 +919,8 @@ class StandaloneImageReferenceTests(TestCase):
|
||||
team=self.team, user=self.user, prompt="不同场景", mode="image", count=1,
|
||||
image_model="volcano", reference_image_ids=[str(ref.id)],
|
||||
)
|
||||
# 任务最终落在支持 image_edit 的模型(gpt-image),而不是用户选的火山 Seedream
|
||||
self.assertIn("gpt-image", tasks[0].model_config.name)
|
||||
self.assertEqual(tasks[0].model_config.provider_id, vp.id)
|
||||
self.assertIn("seedream", tasks[0].model_config.name)
|
||||
|
||||
def test_model_tryon_keeps_each_user_selected_model_even_with_extra_refs(self):
|
||||
"""模特上身图由专用 Worker 向两类模型传商品/人物参考图,不能沿用自由创作的自动换模逻辑。"""
|
||||
|
||||
Reference in New Issue
Block a user