518 lines
20 KiB
Python
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)
|