Files
video-gen/video-gen-api/app/services/model_pricing/rule_service.py
T

518 lines
20 KiB
Python

from __future__ import annotations
import hashlib
import json
from copy import deepcopy
from datetime import datetime, timezone
from typing import Any, Mapping
from sqlalchemy import func, or_, select, text
from sqlalchemy.ext.asyncio import AsyncSession
from app.enums.model_pricing import ModelPricingRuleStatus
from app.models.credit_record import CreditRecord
from app.models.model_pricing_rule import ModelPricingRule
from app.services.model_pricing.calculator import PricingCalculationError, validate_pricing_rule
from app.utils.id_gen import generate_id
PROVIDER_ALIASES = {
"ark": "volcengine",
"volc": "volcengine",
"volc_engine": "volcengine",
"volcano": "volcengine",
"volcengine": "volcengine",
}
class PricingRuleError(ValueError):
pass
def _json_default(value: Any) -> str:
if isinstance(value, datetime):
return value.isoformat()
return str(value)
def canonical_json_hash(value: Mapping[str, Any]) -> str:
raw = json.dumps(value, ensure_ascii=False, sort_keys=True, separators=(",", ":"), default=_json_default)
return hashlib.sha256(raw.encode("utf-8")).hexdigest()
def build_rule_content_hash(
*,
model_category: str,
billing_mode: str,
calculator_version: str,
currency: str,
rule_schema_version: int,
rule_json: Mapping[str, Any],
) -> str:
"""规则正文哈希包含所有会改变计算结果的字段,不只哈希 rule_json。"""
return canonical_json_hash(
{
"model_category": model_category,
"billing_mode": billing_mode,
"calculator_version": calculator_version,
"currency": str(currency or "CNY").upper(),
"rule_schema_version": int(rule_schema_version or 1),
"rule_json": normalize_rule_json(rule_json),
}
)
def normalize_rule_json(value: Mapping[str, Any] | None) -> dict[str, Any]:
"""返回全新的普通 dict,所有 ORM 更新必须整体赋值,禁止嵌套原地修改。"""
return deepcopy(dict(value or {}))
def normalize_provider(provider: str | None, model_name: str | None = None) -> str:
value = str(provider or "").strip().lower()
normalized = PROVIDER_ALIASES.get(value, value)
if normalized in {"sdk", "openai_compatible"} and str(model_name or "").strip().lower().startswith("doubao-"):
return "volcengine"
return normalized
def ensure_aware(value: datetime | None) -> datetime:
value = value or datetime.now(timezone.utc)
return value if value.tzinfo else value.replace(tzinfo=timezone.utc)
def _validate_category_mode(model_category: str, billing_mode: str) -> None:
expected = {
"text_token_tiered": "text",
"image_per_output": "image",
"image_input_output_tiered": "image",
"video_token_rate": "video",
}.get(billing_mode)
if expected is None or model_category != expected:
raise PricingRuleError("模型类型与计价模式不匹配")
def _validate_rule_payload(
*,
model_category: str,
billing_mode: str,
calculator_version: str,
rule_json: Mapping[str, Any],
) -> None:
_validate_category_mode(model_category, billing_mode)
try:
validate_pricing_rule(
billing_mode=billing_mode,
calculator_version=calculator_version,
rule_json=rule_json,
)
except PricingCalculationError as exc:
raise PricingRuleError(str(exc)) from exc
def rule_to_dict(rule: ModelPricingRule) -> dict[str, Any]:
return {
"id": rule.id,
"provider": rule.provider,
"model_name": rule.model_name,
"model_category": rule.model_category,
"billing_mode": rule.billing_mode,
"calculator_version": rule.calculator_version,
"version_code": rule.version_code,
"effective_from": rule.effective_from,
"effective_to": rule.effective_to,
"publish_status": rule.publish_status,
"currency": rule.currency,
"rule_schema_version": rule.rule_schema_version,
"rule_json": deepcopy(rule.rule_json or {}),
"rule_content_hash": rule.rule_content_hash,
"source_url": rule.source_url,
"source_updated_at": rule.source_updated_at,
"remark": rule.remark,
"created_by": rule.created_by,
"updated_by": rule.updated_by,
"created_at": rule.created_at,
"updated_at": rule.updated_at,
}
def _rule_snapshot_query(rule_id: str):
return (
select(
ModelPricingRule.id,
ModelPricingRule.provider,
ModelPricingRule.model_name,
ModelPricingRule.model_category,
ModelPricingRule.billing_mode,
ModelPricingRule.calculator_version,
ModelPricingRule.version_code,
ModelPricingRule.effective_from,
ModelPricingRule.effective_to,
ModelPricingRule.publish_status,
ModelPricingRule.currency,
ModelPricingRule.rule_schema_version,
ModelPricingRule.rule_json,
ModelPricingRule.rule_content_hash,
ModelPricingRule.source_url,
ModelPricingRule.source_updated_at,
ModelPricingRule.remark,
ModelPricingRule.created_by,
ModelPricingRule.updated_by,
ModelPricingRule.created_at,
ModelPricingRule.updated_at,
)
.where(ModelPricingRule.id == rule_id)
.limit(1)
)
def _snapshot_from_mapping(row: Mapping[str, Any]) -> dict[str, Any]:
snapshot = dict(row)
snapshot["rule_json"] = deepcopy(snapshot.get("rule_json") or {})
return snapshot
async def get_rule_snapshot(db: AsyncSession, rule_id: str) -> dict[str, Any]:
"""显式查询并返回普通字典,避免写入后访问过期 ORM 字段触发隐式 IO。"""
row = (await db.execute(_rule_snapshot_query(rule_id))).mappings().one_or_none()
if row is None:
raise PricingRuleError("模型计价规则不存在")
return _snapshot_from_mapping(row)
async def _lock_model_rule_namespace(db: AsyncSession, provider: str, model_name: str) -> None:
bind = db.get_bind()
if bind is not None and bind.dialect.name == "postgresql":
lock_key = f"model_pricing:{provider}:{model_name}"
await db.execute(text("SELECT pg_advisory_xact_lock(hashtext(:lock_key))"), {"lock_key": lock_key})
async def _assert_unique_version(
db: AsyncSession,
*,
provider: str,
model_name: str,
version_code: str,
exclude_id: str | None = None,
) -> None:
query = select(ModelPricingRule.id).where(
ModelPricingRule.provider == provider,
ModelPricingRule.model_name == model_name,
ModelPricingRule.version_code == version_code,
)
if exclude_id:
query = query.where(ModelPricingRule.id != exclude_id)
if (await db.execute(query.limit(1))).scalar_one_or_none():
raise PricingRuleError("该供应商、模型和价格版本号已存在")
async def _assert_no_overlap(
db: AsyncSession,
*,
provider: str,
model_name: str,
effective_from: datetime,
effective_to: datetime | None,
exclude_id: str | None = None,
) -> None:
query = (
select(ModelPricingRule.id)
.where(ModelPricingRule.provider == provider)
.where(ModelPricingRule.model_name == model_name)
.where(ModelPricingRule.publish_status == ModelPricingRuleStatus.PUBLISHED.value)
.where(or_(ModelPricingRule.effective_to.is_(None), ModelPricingRule.effective_to > effective_from))
)
if effective_to is not None:
query = query.where(ModelPricingRule.effective_from < effective_to)
if exclude_id:
query = query.where(ModelPricingRule.id != exclude_id)
if (await db.execute(query.limit(1))).scalar_one_or_none():
raise PricingRuleError("该模型已存在生效时间重叠的已发布价格版本")
async def resolve_published_rule(
db: AsyncSession,
*,
provider: str | None,
model_name: str | None,
reference_at: datetime | None = None,
) -> ModelPricingRule | None:
model_name = str(model_name or "").strip()
provider = normalize_provider(provider, model_name)
if not provider or not model_name:
return None
at = ensure_aware(reference_at)
result = await db.execute(
select(ModelPricingRule)
.where(ModelPricingRule.provider == provider)
.where(ModelPricingRule.model_name == model_name)
.where(ModelPricingRule.publish_status.in_([ModelPricingRuleStatus.PUBLISHED.value, ModelPricingRuleStatus.DISABLED.value]))
.where(ModelPricingRule.effective_from <= at)
.where(or_(ModelPricingRule.effective_to.is_(None), ModelPricingRule.effective_to > at))
.order_by(ModelPricingRule.effective_from.desc(), ModelPricingRule.created_at.desc())
.limit(1)
)
return result.scalar_one_or_none()
async def list_rules(
db: AsyncSession,
*,
page: int = 1,
page_size: int = 50,
provider: str | None = None,
model_name: str | None = None,
model_category: str | None = None,
publish_status: str | None = None,
) -> dict[str, Any]:
filters = []
if provider:
filters.append(ModelPricingRule.provider == normalize_provider(provider))
if model_name:
filters.append(ModelPricingRule.model_name.ilike(f"%{model_name.strip()}%"))
if model_category:
filters.append(ModelPricingRule.model_category == model_category)
if publish_status:
filters.append(ModelPricingRule.publish_status == publish_status)
total = (await db.execute(select(func.count(ModelPricingRule.id)).where(*filters))).scalar_one()
rows = (
await db.execute(
select(ModelPricingRule)
.where(*filters)
.order_by(ModelPricingRule.model_name, ModelPricingRule.effective_from.desc())
.offset((page - 1) * page_size)
.limit(page_size)
)
).scalars().all()
rule_ids = [row.id for row in rows]
referenced: dict[str, int] = {}
if rule_ids:
ref_rows = (
await db.execute(
select(CreditRecord.pricing_rule_id, func.count(CreditRecord.id))
.where(CreditRecord.pricing_rule_id.in_(rule_ids))
.group_by(CreditRecord.pricing_rule_id)
)
).all()
referenced = {str(rule_id): int(count) for rule_id, count in ref_rows if rule_id}
items = []
for row in rows:
item = rule_to_dict(row)
item["referenced_count"] = referenced.get(row.id, 0)
items.append(item)
return {"items": items, "total": int(total or 0)}
async def get_rule(db: AsyncSession, rule_id: str, *, for_update: bool = False) -> ModelPricingRule:
query = select(ModelPricingRule).where(ModelPricingRule.id == rule_id).limit(1)
if for_update:
query = query.with_for_update()
rule = (await db.execute(query)).scalar_one_or_none()
if not rule:
raise PricingRuleError("模型计价规则不存在")
return rule
async def create_rule(db: AsyncSession, *, payload: dict[str, Any], operator_id: str | None) -> dict[str, Any]:
effective_from = ensure_aware(payload["effective_from"])
effective_to = ensure_aware(payload["effective_to"]) if payload.get("effective_to") else None
if effective_to and effective_to <= effective_from:
raise PricingRuleError("失效时间必须晚于生效时间")
model_name = str(payload.get("model_name") or "").strip()
provider = normalize_provider(payload.get("provider"), model_name)
version_code = str(payload.get("version_code") or "").strip()
calculator_version = str(payload.get("calculator_version") or "").strip()
if not provider or not model_name or not version_code or not calculator_version:
raise PricingRuleError("供应商、模型名称、版本号和计算器版本不能为空")
rule_json = normalize_rule_json(payload.get("rule_json"))
_validate_rule_payload(
model_category=payload["model_category"],
billing_mode=payload["billing_mode"],
calculator_version=calculator_version,
rule_json=rule_json,
)
await _lock_model_rule_namespace(db, provider, model_name)
await _assert_unique_version(db, provider=provider, model_name=model_name, version_code=version_code)
rule_id = generate_id()
rule = ModelPricingRule(
id=rule_id,
provider=provider,
model_name=model_name,
model_category=payload["model_category"],
billing_mode=payload["billing_mode"],
calculator_version=calculator_version,
version_code=version_code,
effective_from=effective_from,
effective_to=effective_to,
publish_status=ModelPricingRuleStatus.DRAFT.value,
currency=str(payload.get("currency") or "CNY").upper(),
rule_schema_version=int(payload.get("rule_schema_version") or 1),
rule_json=rule_json,
rule_content_hash=build_rule_content_hash(
model_category=payload["model_category"],
billing_mode=payload["billing_mode"],
calculator_version=calculator_version,
currency=str(payload.get("currency") or "CNY").upper(),
rule_schema_version=int(payload.get("rule_schema_version") or 1),
rule_json=rule_json,
),
source_url=payload.get("source_url"),
source_updated_at=payload.get("source_updated_at"),
remark=payload.get("remark"),
created_by=operator_id,
updated_by=operator_id,
)
db.add(rule)
await db.flush()
return await get_rule_snapshot(db, rule_id)
async def update_draft_rule(
db: AsyncSession,
*,
rule_id: str,
payload: dict[str, Any],
operator_id: str | None,
) -> dict[str, Any]:
rule = await get_rule(db, rule_id, for_update=True)
if rule.publish_status != ModelPricingRuleStatus.DRAFT.value:
raise PricingRuleError("已发布或已停用的价格版本不可修改,请克隆为新版本")
provider = normalize_provider(payload.get("provider", rule.provider), payload.get("model_name", rule.model_name))
model_name = str(payload.get("model_name", rule.model_name) or "").strip()
version_code = str(payload.get("version_code", rule.version_code) or "").strip()
model_category = str(payload.get("model_category", rule.model_category))
billing_mode = str(payload.get("billing_mode", rule.billing_mode))
calculator_version = str(payload.get("calculator_version", rule.calculator_version))
rule_json = normalize_rule_json(payload["rule_json"] if "rule_json" in payload else rule.rule_json)
effective_from = ensure_aware(payload.get("effective_from", rule.effective_from))
effective_to = ensure_aware(payload["effective_to"]) if payload.get("effective_to") else None if "effective_to" in payload else rule.effective_to
if effective_to and effective_to <= effective_from:
raise PricingRuleError("失效时间必须晚于生效时间")
if not provider or not model_name or not version_code or not calculator_version:
raise PricingRuleError("供应商、模型名称、版本号和计算器版本不能为空")
_validate_rule_payload(
model_category=model_category,
billing_mode=billing_mode,
calculator_version=calculator_version,
rule_json=rule_json,
)
await _lock_model_rule_namespace(db, provider, model_name)
await _assert_unique_version(
db,
provider=provider,
model_name=model_name,
version_code=version_code,
exclude_id=rule.id,
)
rule.provider = provider
rule.model_name = model_name
rule.model_category = model_category
rule.billing_mode = billing_mode
rule.calculator_version = calculator_version
rule.version_code = version_code
rule.effective_from = effective_from
rule.effective_to = effective_to
rule.currency = str(payload.get("currency", rule.currency) or "CNY").upper()
rule.rule_schema_version = int(payload.get("rule_schema_version", rule.rule_schema_version) or 1)
rule.rule_json = rule_json
rule.rule_content_hash = build_rule_content_hash(
model_category=rule.model_category,
billing_mode=rule.billing_mode,
calculator_version=rule.calculator_version,
currency=rule.currency,
rule_schema_version=rule.rule_schema_version,
rule_json=rule.rule_json,
)
for key in ("source_url", "source_updated_at", "remark"):
if key in payload:
setattr(rule, key, payload[key])
rule.updated_by = operator_id
await db.flush()
return await get_rule_snapshot(db, rule_id)
async def publish_rule(db: AsyncSession, *, rule_id: str, operator_id: str | None) -> dict[str, Any]:
rule = await get_rule(db, rule_id, for_update=True)
publish_status = rule.publish_status
if publish_status == ModelPricingRuleStatus.PUBLISHED.value:
return await get_rule_snapshot(db, rule_id)
if publish_status != ModelPricingRuleStatus.DRAFT.value:
raise PricingRuleError("只有草稿价格版本可以发布")
provider = rule.provider
model_name = rule.model_name
model_category = rule.model_category
billing_mode = rule.billing_mode
calculator_version = rule.calculator_version
effective_from = rule.effective_from
effective_to = rule.effective_to
currency = rule.currency
rule_schema_version = rule.rule_schema_version
rule_json = normalize_rule_json(rule.rule_json)
await _lock_model_rule_namespace(db, provider, model_name)
_validate_rule_payload(
model_category=model_category,
billing_mode=billing_mode,
calculator_version=calculator_version,
rule_json=rule_json,
)
previous = (
await db.execute(
select(ModelPricingRule)
.where(ModelPricingRule.provider == provider)
.where(ModelPricingRule.model_name == model_name)
.where(ModelPricingRule.publish_status == ModelPricingRuleStatus.PUBLISHED.value)
.where(ModelPricingRule.effective_from < effective_from)
.where(ModelPricingRule.effective_to.is_(None))
.order_by(ModelPricingRule.effective_from.desc())
.limit(1)
.with_for_update()
)
).scalar_one_or_none()
if previous:
previous.effective_to = effective_from
previous.updated_by = operator_id
await db.flush()
await _assert_no_overlap(
db,
provider=provider,
model_name=model_name,
effective_from=effective_from,
effective_to=effective_to,
exclude_id=rule_id,
)
rule.rule_json = rule_json
rule.rule_content_hash = build_rule_content_hash(
model_category=model_category,
billing_mode=billing_mode,
calculator_version=calculator_version,
currency=currency,
rule_schema_version=rule_schema_version,
rule_json=rule_json,
)
rule.publish_status = ModelPricingRuleStatus.PUBLISHED.value
rule.updated_by = operator_id
await db.flush()
return await get_rule_snapshot(db, rule_id)
async def disable_rule(db: AsyncSession, *, rule_id: str, operator_id: str | None) -> dict[str, Any]:
rule = await get_rule(db, rule_id, for_update=True)
if rule.publish_status == ModelPricingRuleStatus.DISABLED.value:
return await get_rule_snapshot(db, rule_id)
await _lock_model_rule_namespace(db, rule.provider, rule.model_name)
if rule.publish_status == ModelPricingRuleStatus.PUBLISHED.value:
now = datetime.now(timezone.utc)
close_at = now if rule.effective_from < now else rule.effective_from
if rule.effective_to is None or rule.effective_to > close_at:
rule.effective_to = close_at
rule.publish_status = ModelPricingRuleStatus.DISABLED.value
rule.updated_by = operator_id
await db.flush()
return await get_rule_snapshot(db, rule_id)