feat: 接入模型动态 Fallback 与调用审计
This commit is contained in:
@@ -3,7 +3,8 @@ from decimal import Decimal
|
||||
from rest_framework import serializers
|
||||
|
||||
from apps.accounts.models import Team, TeamMember, User
|
||||
from apps.ai.models import AITask, ModelConfig, ModelProvider, PromptTemplate, QualityWord
|
||||
from apps.ai.model_routing import model_metadata_errors, provider_metadata_errors
|
||||
from apps.ai.models import AIModelAttempt, AITask, ModelConfig, ModelProvider, PromptTemplate, QualityWord
|
||||
from apps.assets.models import Asset
|
||||
from apps.billing.models import CreditLedger, QuotaPolicy
|
||||
from apps.projects.models import Project
|
||||
@@ -162,11 +163,26 @@ class AdminTaskSerializer(serializers.ModelSerializer):
|
||||
return str((actual / rate - base).quantize(Decimal("0.01")))
|
||||
|
||||
|
||||
class AdminModelAttemptSerializer(serializers.ModelSerializer):
|
||||
class Meta:
|
||||
model = AIModelAttempt
|
||||
fields = [
|
||||
"id", "sequence", "provider_name", "provider_display_name", "model_name", "model_display_name",
|
||||
"public_model_name", "capability", "operation", "status", "is_retry", "is_fallback",
|
||||
"previous_attempt", "provider_task_id", "started_at", "finished_at", "duration_ms", "error_type",
|
||||
"provider_error_code", "raw_error", "safe_error_summary", "usage", "platform_cost", "request_summary",
|
||||
"response_summary",
|
||||
]
|
||||
read_only_fields = fields
|
||||
|
||||
|
||||
class AdminTaskDetailSerializer(AdminTaskSerializer):
|
||||
attempts = AdminModelAttemptSerializer(source="model_attempts", many=True, read_only=True)
|
||||
|
||||
class Meta(AdminTaskSerializer.Meta):
|
||||
fields = AdminTaskSerializer.Meta.fields + [
|
||||
"project", "idempotency_key", "request_payload", "response_payload",
|
||||
"error_message", "submitted_at", "completed_at",
|
||||
"error_message", "submitted_at", "completed_at", "attempts",
|
||||
]
|
||||
read_only_fields = fields
|
||||
|
||||
@@ -211,6 +227,12 @@ class AdminModelProviderSerializer(serializers.ModelSerializer):
|
||||
anno = getattr(obj, "model_count_anno", None)
|
||||
return anno if anno is not None else obj.models.count()
|
||||
|
||||
def validate_metadata(self, value):
|
||||
errors = provider_metadata_errors(value)
|
||||
if errors:
|
||||
raise serializers.ValidationError(list(errors))
|
||||
return value
|
||||
|
||||
|
||||
class AdminModelConfigSerializer(serializers.ModelSerializer):
|
||||
provider_name = serializers.CharField(source="provider.name", read_only=True, default=None)
|
||||
@@ -224,6 +246,16 @@ class AdminModelConfigSerializer(serializers.ModelSerializer):
|
||||
# is_default 只经 set-default 端点改,不在普通编辑里直接写
|
||||
read_only_fields = ["id", "provider_name", "is_default", "created_at"]
|
||||
|
||||
def validate(self, attrs):
|
||||
attrs = super().validate(attrs)
|
||||
if self.instance is None or "metadata" in attrs or "capability" in attrs:
|
||||
capability = attrs.get("capability", getattr(self.instance, "capability", ""))
|
||||
metadata = attrs.get("metadata", getattr(self.instance, "metadata", {}))
|
||||
errors = model_metadata_errors(capability, metadata)
|
||||
if errors:
|
||||
raise serializers.ValidationError({"metadata": list(errors)})
|
||||
return attrs
|
||||
|
||||
|
||||
class AdminProjectSerializer(serializers.ModelSerializer):
|
||||
team_name = serializers.CharField(source="team.name", read_only=True, default=None)
|
||||
|
||||
@@ -368,6 +368,35 @@ class AdminTaskMonitorTests(TestCase):
|
||||
self.assertEqual(r.status_code, 200)
|
||||
self.assertGreaterEqual(r.data["count"], 4)
|
||||
|
||||
def test_list_defers_large_payload_columns_but_detail_keeps_them(self):
|
||||
from django.db import connection
|
||||
from django.test.utils import CaptureQueriesContext
|
||||
|
||||
self.t_ok.request_payload = {"prompt": "x" * 100_000}
|
||||
self.t_ok.response_payload = {"image": "y" * 100_000}
|
||||
self.t_ok.error_message = "z" * 100_000
|
||||
self.t_ok.save(update_fields=["request_payload", "response_payload", "error_message", "updated_at"])
|
||||
|
||||
with CaptureQueriesContext(connection) as captured:
|
||||
response = self.ac.get("/api/admin/tasks/?page_size=10")
|
||||
self.assertEqual(response.status_code, 200)
|
||||
task_selects = [
|
||||
item["sql"]
|
||||
for item in captured.captured_queries
|
||||
if "FROM \"ai_aitask\"" in item["sql"] and "COUNT(" not in item["sql"]
|
||||
]
|
||||
self.assertTrue(task_selects)
|
||||
for sql in task_selects:
|
||||
self.assertNotIn("request_payload", sql)
|
||||
self.assertNotIn("response_payload", sql)
|
||||
self.assertNotIn("error_message", sql)
|
||||
|
||||
detail = self.ac.get(f"/api/admin/tasks/{self.t_ok.id}/")
|
||||
self.assertEqual(detail.status_code, 200)
|
||||
self.assertEqual(detail.data["request_payload"], self.t_ok.request_payload)
|
||||
self.assertEqual(detail.data["response_payload"], self.t_ok.response_payload)
|
||||
self.assertEqual(detail.data["error_message"], self.t_ok.error_message)
|
||||
|
||||
def test_filter_status_type_anomaly(self):
|
||||
self.assertTrue(all(t["status"] == "failed" for t in self.ac.get("/api/admin/tasks/?status=failed").data["results"]))
|
||||
self.assertTrue(all(t["task_type"] == "person_image" for t in self.ac.get("/api/admin/tasks/?task_type=person_image").data["results"]))
|
||||
@@ -380,10 +409,34 @@ class AdminTaskMonitorTests(TestCase):
|
||||
self.assertFalse(self.ac.get(f"/api/admin/tasks/{self.t_ok.id}/").data["cost_anomaly"])
|
||||
|
||||
def test_detail_has_payloads(self):
|
||||
from django.utils import timezone
|
||||
|
||||
from apps.ai.models import AIModelAttempt
|
||||
|
||||
AIModelAttempt.objects.create(
|
||||
task=self.t_failed,
|
||||
sequence=1,
|
||||
provider=self.mc.provider,
|
||||
model_config=self.mc,
|
||||
provider_name=self.mc.provider.name,
|
||||
provider_display_name=self.mc.provider.display_name,
|
||||
model_name=self.mc.name,
|
||||
model_display_name=self.mc.display_name,
|
||||
public_model_name="AirShelf Image",
|
||||
capability="image",
|
||||
operation="image_generate",
|
||||
status=AIModelAttempt.Status.FAILED,
|
||||
started_at=timezone.now(),
|
||||
error_type="provider_unavailable",
|
||||
)
|
||||
d = self.ac.get(f"/api/admin/tasks/{self.t_failed.id}/")
|
||||
self.assertEqual(d.status_code, 200)
|
||||
self.assertIn("request_payload", d.data)
|
||||
self.assertIn("response_payload", d.data)
|
||||
self.assertEqual(len(d.data["attempts"]), 1)
|
||||
self.assertEqual(d.data["attempts"][0]["model_name"], self.mc.name)
|
||||
row = next(item for item in self.ac.get("/api/admin/tasks/").data["results"] if item["id"] == str(self.t_failed.id))
|
||||
self.assertNotIn("attempts", row)
|
||||
|
||||
def test_retry_dispatch_and_audit(self):
|
||||
with patch("apps.ai.tasks.generate_standalone_image_task.delay") as mock_delay:
|
||||
|
||||
@@ -417,7 +417,13 @@ def admin_asset_reviews_poll(request):
|
||||
@permission_classes([IsPlatformAdmin])
|
||||
def admin_tasks(request):
|
||||
"""全局 AITask 列表(?status= / ?task_type= / ?team= / ?anomaly=1 成本异常 筛 + 分页)。"""
|
||||
qs = AITask.objects.select_related("team", "model_config").order_by("-created_at")
|
||||
# 列表不返回请求/响应/完整错误正文。部分图片、视频任务的 JSON 可达数 MB,若随列表页
|
||||
# 从远程 MySQL 读取,会让只有 10 行的分页请求也长时间卡在“加载中”;详情接口仍完整读取。
|
||||
qs = (
|
||||
AITask.objects.select_related("team", "model_config")
|
||||
.defer("request_payload", "response_payload", "error_message")
|
||||
.order_by("-created_at")
|
||||
)
|
||||
st = request.query_params.get("status")
|
||||
if st in dict(AITask.Status.choices):
|
||||
qs = qs.filter(status=st)
|
||||
@@ -437,7 +443,13 @@ def admin_tasks(request):
|
||||
@api_view(["GET"])
|
||||
@permission_classes([IsPlatformAdmin])
|
||||
def admin_task_detail(request, task_id):
|
||||
task = AITask.objects.select_related("team", "model_config").filter(id=task_id).first()
|
||||
# 尝试链只在详情请求加载,列表仍保持原查询与一任务一行。
|
||||
task = (
|
||||
AITask.objects.select_related("team", "model_config")
|
||||
.prefetch_related("model_attempts")
|
||||
.filter(id=task_id)
|
||||
.first()
|
||||
)
|
||||
if task is None:
|
||||
return Response({"detail": "not found"}, status=status.HTTP_404_NOT_FOUND)
|
||||
return Response(AdminTaskDetailSerializer(task).data)
|
||||
|
||||
Reference in New Issue
Block a user