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)