333 lines
14 KiB
Python
333 lines
14 KiB
Python
from decimal import Decimal
|
||
|
||
from rest_framework import serializers
|
||
|
||
from apps.accounts.models import Team, TeamMember, User
|
||
from apps.ai.model_routing import model_metadata_errors, provider_metadata_errors
|
||
from apps.billing.points_rules import normalize_points_pricing_for_save, points_pricing_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
|
||
|
||
# 成本异常阈值:实际成本 > 预估 × 此倍数(且预估 > 0)即标异常
|
||
COST_ANOMALY_RATIO = Decimal("1.5")
|
||
|
||
|
||
def is_cost_anomaly(estimated, actual) -> bool:
|
||
est = estimated or Decimal("0")
|
||
act = actual or Decimal("0")
|
||
return bool(est > 0 and act > est * COST_ANOMALY_RATIO)
|
||
|
||
|
||
class QualityWordSerializer(serializers.ModelSerializer):
|
||
class Meta:
|
||
model = QualityWord
|
||
fields = ["id", "stage", "slot", "text", "sort", "enabled", "created_at"]
|
||
read_only_fields = ["id", "created_at"]
|
||
|
||
|
||
# 每条提示词可用的 {占位符}(给 admin 提示;运行时由对应 builder 注入,改错会原样保留不致命)
|
||
PROMPT_PLACEHOLDERS = {
|
||
"person_portrait": ["描述"],
|
||
"person_triview": [],
|
||
"product_triview": ["商品", "补充"],
|
||
"scene": ["场景描述"],
|
||
"video_segment": ["开场", "设定", "风格", "脚本", "时长"],
|
||
}
|
||
|
||
|
||
class PromptTemplateSerializer(serializers.ModelSerializer):
|
||
label = serializers.CharField(source="get_key_display", read_only=True)
|
||
placeholders = serializers.SerializerMethodField()
|
||
|
||
class Meta:
|
||
model = PromptTemplate
|
||
fields = ["id", "key", "label", "template", "ratio", "enabled", "placeholders", "updated_at"]
|
||
read_only_fields = ["id", "key", "label", "placeholders", "updated_at"] # key 固定,只改正文/比例/启用
|
||
|
||
def get_placeholders(self, obj) -> list:
|
||
return PROMPT_PLACEHOLDERS.get(obj.key, [])
|
||
|
||
|
||
class AdminReviewAssetSerializer(serializers.ModelSerializer):
|
||
"""跨团队人像审核队列行:复用 AssetFileSerializer 的首图 preview_url。"""
|
||
|
||
team_name = serializers.CharField(source="team.name", read_only=True, default=None)
|
||
preview_url = serializers.SerializerMethodField()
|
||
|
||
class Meta:
|
||
model = Asset
|
||
fields = ["id", "name", "category", "review_status", "review_error", "team", "team_name", "preview_url", "created_at"]
|
||
read_only_fields = fields
|
||
|
||
def get_preview_url(self, obj):
|
||
from apps.assets.serializers import AssetFileSerializer
|
||
|
||
# 用 prefetch 缓存(list(...))而非 .first():后者会另发查询,列表页 N+1 拖慢到数秒
|
||
files = list(obj.files.all())
|
||
return AssetFileSerializer(files[0]).data.get("preview_url", "") if files else ""
|
||
|
||
|
||
class AdminTeamSerializer(serializers.ModelSerializer):
|
||
owner_username = serializers.CharField(source="owner.username", read_only=True, default=None)
|
||
member_count = serializers.SerializerMethodField()
|
||
balance = serializers.SerializerMethodField()
|
||
|
||
class Meta:
|
||
model = Team
|
||
fields = ["id", "name", "status", "owner", "owner_username", "member_count", "balance", "price_multiplier", "created_at"]
|
||
read_only_fields = fields
|
||
|
||
def get_member_count(self, obj):
|
||
anno = getattr(obj, "member_count_anno", None)
|
||
return anno if anno is not None else obj.members.count()
|
||
|
||
def get_balance(self, obj):
|
||
acct = getattr(obj, "credit_account", None)
|
||
return str(acct.balance) if acct is not None else "0"
|
||
|
||
|
||
class AdminTeamMemberSerializer(serializers.ModelSerializer):
|
||
username = serializers.CharField(source="user.username", read_only=True)
|
||
user_status = serializers.CharField(source="user.status", read_only=True)
|
||
|
||
class Meta:
|
||
model = TeamMember
|
||
fields = ["id", "username", "role", "status", "user_status", "monthly_credit_limit"]
|
||
read_only_fields = fields
|
||
|
||
|
||
def wallet_membership(user):
|
||
"""用户的钱包归属:第一个 active 成员关系(与 common.api.get_current_team 同口径,消费从这个团队扣)。
|
||
列表页依赖 prefetch 好的 team_memberships,这里只在内存里挑,不要改成 .filter() 否则每行多一次查询。"""
|
||
actives = [m for m in user.team_memberships.all() if m.status == TeamMember.Status.ACTIVE]
|
||
return min(actives, key=lambda m: m.created_at) if actives else None
|
||
|
||
|
||
class AdminUserSerializer(serializers.ModelSerializer):
|
||
teams = serializers.SerializerMethodField()
|
||
# 后台按「用户」视角管积分,但账户实际挂在团队上。这里把钱包团队拍平给前端:
|
||
# wallet_shared=True 表示这是个多人共享池(老用户),发积分弹窗要明确提示钱会进整个团队。
|
||
balance = serializers.SerializerMethodField()
|
||
wallet_team = serializers.SerializerMethodField()
|
||
wallet_team_name = serializers.SerializerMethodField()
|
||
wallet_shared = serializers.SerializerMethodField()
|
||
|
||
class Meta:
|
||
model = User
|
||
fields = [
|
||
"id", "username", "status", "is_platform_admin", "date_joined", "teams",
|
||
"balance", "wallet_team", "wallet_team_name", "wallet_shared",
|
||
]
|
||
read_only_fields = fields
|
||
|
||
def get_teams(self, obj):
|
||
return [
|
||
{"team_id": str(m.team_id), "team_name": m.team.name, "role": m.role}
|
||
for m in obj.team_memberships.all()
|
||
if not m.team.is_personal # 个人团队是钱包实现细节,不当「所属团队」展示
|
||
]
|
||
|
||
def get_balance(self, obj):
|
||
m = wallet_membership(obj)
|
||
acct = getattr(m.team, "credit_account", None) if m else None
|
||
return str(acct.balance) if acct is not None else "0"
|
||
|
||
def get_wallet_team(self, obj):
|
||
m = wallet_membership(obj)
|
||
return str(m.team_id) if m else None
|
||
|
||
def get_wallet_team_name(self, obj):
|
||
m = wallet_membership(obj)
|
||
return m.team.name if m else None
|
||
|
||
def get_wallet_shared(self, obj):
|
||
m = wallet_membership(obj)
|
||
if m is None:
|
||
return False
|
||
count = getattr(m, "team_member_count", None)
|
||
return bool((count if count is not None else m.team.members.count()) > 1)
|
||
|
||
|
||
class AdminTaskSerializer(serializers.ModelSerializer):
|
||
team_name = serializers.CharField(source="team.name", read_only=True, default=None)
|
||
model_name = serializers.CharField(source="model_config.name", read_only=True, default=None)
|
||
# 不改变 AITask.task_type 的调度语义;后台展示/筛选使用来源分类。
|
||
task_category = serializers.SerializerMethodField()
|
||
cost_anomaly = serializers.SerializerMethodField()
|
||
# 单任务毛利(¥):actual_cost(积分)÷汇率 − base_cost。base_cost=0(成本未知)时 None,报表侧过滤
|
||
margin_yuan = serializers.SerializerMethodField()
|
||
# 可手动回收:卡在 RESERVED 超过 10 分钟(与 admin_task_reap 的服务端闸完全同口径,前端据此显示按钮)
|
||
reapable = serializers.SerializerMethodField()
|
||
|
||
class Meta:
|
||
model = AITask
|
||
fields = [
|
||
"id", "task_type", "task_category", "status", "team", "team_name", "model_name",
|
||
"estimated_cost", "actual_cost", "base_cost", "margin_yuan", "cost_anomaly", "error_code", "reapable", "created_at",
|
||
]
|
||
read_only_fields = fields
|
||
|
||
def get_task_category(self, obj) -> str:
|
||
return str(getattr(obj, "task_category", "standard") or "standard")
|
||
|
||
def get_cost_anomaly(self, obj) -> bool:
|
||
return is_cost_anomaly(obj.estimated_cost, obj.actual_cost)
|
||
|
||
def get_reapable(self, obj) -> bool:
|
||
from datetime import timedelta
|
||
|
||
from django.utils import timezone
|
||
|
||
return bool(
|
||
obj.status == AITask.Status.RESERVED
|
||
and obj.updated_at is not None
|
||
and obj.updated_at < timezone.now() - timedelta(minutes=10)
|
||
)
|
||
|
||
def get_margin_yuan(self, obj) -> str | None:
|
||
base = obj.base_cost or Decimal("0")
|
||
actual = obj.actual_cost or Decimal("0")
|
||
if base <= 0 or actual <= 0:
|
||
return None
|
||
from apps.billing.pricing import get_billing_config
|
||
|
||
# 优先用任务计价当时的汇率快照:汇率调整不追溯历史,毛利报表不整体漂移(review 确认)。
|
||
# __dict__ 直取避免触发 deferred 列加载(有的列表 queryset 会 defer payload)。
|
||
payload = obj.__dict__.get("request_payload") or {}
|
||
try:
|
||
rate = Decimal(str(payload.get("points_per_yuan_snapshot") or "")) if payload.get("points_per_yuan_snapshot") else get_billing_config().points_per_yuan
|
||
except Exception: # noqa: BLE001
|
||
rate = get_billing_config().points_per_yuan
|
||
if rate <= 0:
|
||
return None
|
||
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", "attempts",
|
||
]
|
||
read_only_fields = fields
|
||
|
||
|
||
class AdminLedgerSerializer(serializers.ModelSerializer):
|
||
team_name = serializers.CharField(source="team.name", read_only=True, default=None)
|
||
username = serializers.CharField(source="user.username", read_only=True, default=None)
|
||
|
||
class Meta:
|
||
model = CreditLedger
|
||
fields = ["id", "team", "team_name", "username", "ledger_type", "amount", "balance_after", "reason", "created_at"]
|
||
read_only_fields = fields
|
||
|
||
|
||
class AdminQuotaPolicySerializer(serializers.ModelSerializer):
|
||
team_name = serializers.CharField(source="team.name", read_only=True, default=None)
|
||
|
||
class Meta:
|
||
model = QuotaPolicy
|
||
fields = [
|
||
"id", "team", "team_name", "user", "project",
|
||
"monthly_limit", "project_limit", "per_task_limit", "is_active", "created_at",
|
||
]
|
||
read_only_fields = ["id", "team_name", "created_at"]
|
||
|
||
|
||
class AdminModelProviderSerializer(serializers.ModelSerializer):
|
||
has_api_key = serializers.SerializerMethodField()
|
||
model_count = serializers.SerializerMethodField()
|
||
|
||
class Meta:
|
||
model = ModelProvider
|
||
fields = ["id", "name", "display_name", "status", "base_url", "api_key", "metadata", "has_api_key", "model_count", "created_at"]
|
||
# api_key 只写不读(密钥不回传前端)
|
||
extra_kwargs = {"api_key": {"write_only": True, "required": False}}
|
||
read_only_fields = ["id", "has_api_key", "model_count", "created_at"]
|
||
|
||
def get_has_api_key(self, obj) -> bool:
|
||
return bool(obj.api_key)
|
||
|
||
def get_model_count(self, obj) -> int:
|
||
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)
|
||
|
||
class Meta:
|
||
model = ModelConfig
|
||
fields = [
|
||
"id", "provider", "provider_name", "name", "display_name", "capability",
|
||
"endpoint", "unit_price", "status", "is_default", "rate_limit_per_minute", "metadata", "created_at",
|
||
]
|
||
# 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 or "unit_price" in attrs:
|
||
capability = attrs.get("capability", getattr(self.instance, "capability", ""))
|
||
metadata = attrs.get("metadata", getattr(self.instance, "metadata", {}))
|
||
if not isinstance(metadata, dict):
|
||
metadata = {}
|
||
metadata = normalize_points_pricing_for_save(capability, metadata)
|
||
# 图片/文本:表单 unit_price 与挂牌积分同步
|
||
if "unit_price" in attrs and capability in {"image", "text", "vision"}:
|
||
from decimal import Decimal
|
||
try:
|
||
pts = int(Decimal(str(attrs.get("unit_price") or 0)))
|
||
except Exception:
|
||
pts = 0
|
||
pricing = dict(metadata.get("points_pricing") or {})
|
||
if capability == "image":
|
||
pricing.update({"mode": "per_image", "points_per_image": max(pts, 0)})
|
||
else:
|
||
pricing.update({"mode": "per_call", "points_per_call": max(pts, 0)})
|
||
metadata["points_pricing"] = pricing
|
||
attrs["metadata"] = metadata
|
||
# 视频有挂牌秒价时,把展示用 unit_price 写成最低档秒积分,列表更好读
|
||
if capability == "video":
|
||
tiers = (metadata.get("points_pricing") or {}).get("tiers") or []
|
||
secs = [t.get("points_per_second") for t in tiers if isinstance(t, dict) and t.get("points_per_second") is not None]
|
||
if secs:
|
||
attrs["unit_price"] = min(int(x) for x in secs)
|
||
errors = list(model_metadata_errors(capability, metadata)) + list(points_pricing_errors(capability, metadata))
|
||
if errors:
|
||
raise serializers.ValidationError({"metadata": errors})
|
||
return attrs
|
||
|
||
|
||
class AdminProjectSerializer(serializers.ModelSerializer):
|
||
team_name = serializers.CharField(source="team.name", read_only=True, default=None)
|
||
product_title = serializers.CharField(source="product.title", read_only=True, default=None)
|
||
|
||
class Meta:
|
||
model = Project
|
||
fields = ["id", "name", "team", "team_name", "product_title", "status", "current_stage", "created_at"]
|
||
read_only_fields = fields
|