This commit is contained in:
2026-07-10 17:09:14 +08:00
parent ad2e40bc67
commit d07cd508ee
54 changed files with 7111 additions and 619 deletions
@@ -0,0 +1,981 @@
from __future__ import annotations
import argparse
import asyncio
from copy import deepcopy
from datetime import datetime, timezone
from decimal import Decimal
from typing import Any, Mapping
from sqlalchemy import or_, select
from app.enums.credit_record import (
CreditRecordAction,
CreditRecordChargeKind,
CreditRecordOwnerType,
CreditRecordType,
)
from app.enums.model_pricing import (
ModelPricingRuleStatus,
PricingSnapshotStage,
ProviderCostStatus,
)
from app.models.base import async_session
from app.models.chat_generation_task import ChatGenerationTask
from app.models.credit_record import CreditRecord
from app.models.generated_resource import GeneratedResource
from app.models.generation_record import GenerationRecord
from app.models.image_engine import ImageEngine
from app.models.model_config import ModelConfig
from app.models.model_pricing_rule import ModelPricingRule
from app.models.module_generation_project import ModuleGenerationProject
from app.models.module_generation_step import ModuleGenerationStep
from app.models.shot_replicate_segment import ShotReplicateSegment
from app.models.shot_replicate_task_set import ShotReplicateTaskSet
from app.models.token_usage import TokenUsage
from app.models.video_engine import VideoEngine
from app.services.model_pricing.attachment_snapshot_service import (
build_attachment_snapshot,
build_generation_snapshot,
)
from app.services.model_pricing.rule_service import normalize_provider
from app.services.model_pricing.snapshot_service import finalize_credit_record_pricing
from app.services.model_pricing.usage_normalizer import (
normalize_provider_media_usage,
parse_size,
safe_float,
safe_int,
safe_json_dict,
)
from app.services.operation_log_service import log_model_pricing_event
from app.services.resource_accounting_service import (
SOURCE_MODEL_CHAT_TASK,
SOURCE_MODEL_GENERATION_RECORD,
)
PROCESSABLE_CHARGE_KINDS = {
CreditRecordChargeKind.MEDIA.value,
CreditRecordChargeKind.TEXT_PROMPT.value,
CreditRecordChargeKind.VIDEO_ANALYSIS.value,
}
NON_PROVIDER_CHARGE_KINDS = {
CreditRecordChargeKind.FILE_PARSE.value,
CreditRecordChargeKind.VISION_INPUT.value,
CreditRecordChargeKind.MODULE_CREATE.value,
CreditRecordChargeKind.VIDEO_SPLIT.value,
CreditRecordChargeKind.RECHARGE.value,
CreditRecordChargeKind.REFUND.value,
CreditRecordChargeKind.ADMIN_ADJUST.value,
CreditRecordChargeKind.TEAM_INTERNAL.value,
}
INCOMPLETE_COST_STATUSES = {
None,
"",
ProviderCostStatus.PENDING.value,
ProviderCostStatus.UNMATCHED_RULE.value,
ProviderCostStatus.USAGE_MISSING.value,
ProviderCostStatus.ERROR.value,
ProviderCostStatus.PROVIDER_RESULT_UNCERTAIN.value,
ProviderCostStatus.HISTORICAL_PRICE_UNAVAILABLE.value,
ProviderCostStatus.HISTORICAL_ENGINE_UNAVAILABLE.value,
}
class BackfillContext:
def __init__(self) -> None:
self.owner_maps: dict[str, dict[str, Any]] = {}
self.linked_chat_tasks: dict[str, ChatGenerationTask] = {}
self.token_usage_by_id: dict[str, TokenUsage] = {}
self.token_usage_by_owner: dict[tuple[str, str], TokenUsage] = {}
self.model_configs: dict[str, ModelConfig] = {}
self.image_engines: dict[str, ImageEngine] = {}
self.video_engines: dict[str, VideoEngine] = {}
self.resource_counts: dict[tuple[str, str], dict[str, int]] = {}
self.current_rules_by_category: dict[str, list[ModelPricingRule]] = {}
class BackfillStats:
def __init__(self) -> None:
self.scanned = 0
self.changed = 0
self.calculated = 0
self.estimated = 0
self.rule_bound = 0
self.missing_engine = 0
self.missing_usage = 0
self.unmatched_rule = 0
self.not_applicable = 0
self.skipped_completed = 0
self.skipped_non_provider = 0
self.failed = 0
self.force_repriced = 0
def merge(self, other: "BackfillStats") -> None:
for key in vars(self):
setattr(self, key, getattr(self, key) + getattr(other, key))
def as_dict(self) -> dict[str, int]:
return {key: int(value) for key, value in vars(self).items()}
def _parse_date(value: str | None, *, end: bool = False) -> datetime | None:
if not value:
return None
parsed = datetime.fromisoformat(value)
if parsed.tzinfo is None:
parsed = parsed.replace(tzinfo=timezone.utc)
if end and len(value) <= 10:
parsed = parsed.replace(hour=23, minute=59, second=59, microsecond=999999)
return parsed
def _utcnow() -> datetime:
return datetime.now(timezone.utc)
def _owner_key(record: CreditRecord) -> tuple[str, str] | None:
owner_type = str(record.owner_type or "").strip()
owner_id = str(record.owner_id or record.related_id or "").strip()
return (owner_type, owner_id) if owner_type and owner_id else None
def _is_provider_cost_candidate(record: CreditRecord) -> bool:
if record.type != CreditRecordType.CONSUME.value:
return False
if record.charge_action not in {None, "", CreditRecordAction.CHARGE.value}:
return False
charge_kind = str(record.charge_kind or "").strip()
if charge_kind in NON_PROVIDER_CHARGE_KINDS:
return False
if charge_kind in PROCESSABLE_CHARGE_KINDS:
return True
if str(record.media_type or "").lower() in {"image", "video"}:
return True
if any(int(value or 0) > 0 for value in (record.input_tokens, record.output_tokens, record.total_tokens)):
return True
return bool(record.token_usage_id or _owner_key(record))
def _raw_references(owner: Any) -> Any:
if owner is None:
return None
if isinstance(owner, ShotReplicateTaskSet):
return [
{
"type": "video",
"url": owner.video_url,
"path": owner.video_path,
"duration_seconds": owner.video_duration_seconds,
"role": "reference_video",
"billable_input": True,
}
]
if isinstance(owner, ShotReplicateSegment):
return [
{
"type": "video",
"url": owner.segment_video_url,
"path": owner.segment_video_path,
"duration_seconds": owner.duration_seconds,
"role": "reference_video",
"billable_input": True,
}
]
if hasattr(owner, "media_references"):
return getattr(owner, "media_references", None)
if hasattr(owner, "input_json"):
return getattr(owner, "input_json", None)
return None
def _provider_response(owner: Any) -> Any:
if owner is None:
return None
for name in (
"provider_response_json",
"output_json",
"analysis_raw_json",
"analysis_result_json",
"analysis_json",
):
value = getattr(owner, name, None)
if value:
return value
return None
def _nested_usage(value: Any) -> dict[str, Any]:
data = safe_json_dict(value)
if not data:
return {}
usage = data.get("usage")
if isinstance(usage, Mapping):
return deepcopy(dict(usage))
for key in ("result", "payload", "data", "response"):
child = data.get(key)
if isinstance(child, Mapping):
found = _nested_usage(child)
if found:
return found
return {}
def _select_token_usage(record: CreditRecord, owner: Any, ctx: BackfillContext) -> TokenUsage | None:
token_usage_id = (
record.token_usage_id
or getattr(owner, "token_usage_id", None)
)
if token_usage_id and token_usage_id in ctx.token_usage_by_id:
return ctx.token_usage_by_id[token_usage_id]
key = _owner_key(record)
if key and key in ctx.token_usage_by_owner:
return ctx.token_usage_by_owner[key]
return None
def _usage_from_record(record: CreditRecord, owner: Any, token_usage: TokenUsage | None) -> dict[str, Any]:
usage = deepcopy(dict(record.usage_snapshot_json or {}))
payload_usage = _nested_usage(_provider_response(owner))
for key, value in payload_usage.items():
usage.setdefault(key, value)
owner_input = safe_int(getattr(owner, "input_tokens", None))
owner_output = safe_int(getattr(owner, "output_tokens", None))
owner_total = safe_int(getattr(owner, "total_tokens", None))
token_input = safe_int(getattr(token_usage, "input_tokens", None))
token_output = safe_int(getattr(token_usage, "output_tokens", None))
token_total = safe_int(getattr(token_usage, "total_tokens", None))
input_tokens = max(0, safe_int(record.input_tokens, token_input or owner_input))
output_tokens = max(0, safe_int(record.output_tokens, token_output or owner_output))
total_tokens = max(
0,
safe_int(record.total_tokens, token_total or owner_total or (input_tokens + output_tokens)),
)
if total_tokens <= 0:
total_tokens = input_tokens + output_tokens
if output_tokens <= 0 and total_tokens > input_tokens:
output_tokens = total_tokens - input_tokens
usage.update(
{
"input_tokens": input_tokens,
"output_tokens": output_tokens,
"total_tokens": total_tokens,
"context_tokens": max(0, safe_int(usage.get("context_tokens"), input_tokens)),
"usage_source": "backfill",
"provider_usage_primary": bool(record.provider_usage_primary),
}
)
return usage
def _get_engine_object(record: CreditRecord, owner: Any, ctx: BackfillContext) -> Any:
engine_id = str(record.engine_id or getattr(owner, "engine_id", None) or "").strip()
if not engine_id:
return None
media_type = str(record.media_type or getattr(owner, "gen_type", None) or "").lower()
if media_type == "image":
return ctx.image_engines.get(engine_id)
if media_type == "video":
return ctx.video_engines.get(engine_id)
return ctx.image_engines.get(engine_id) or ctx.video_engines.get(engine_id)
def _restore_engine_snapshot(
record: CreditRecord,
*,
owner: Any,
linked_chat: ChatGenerationTask | None,
token_usage: TokenUsage | None,
ctx: BackfillContext,
) -> None:
response = safe_json_dict(_provider_response(linked_chat or owner))
owner_snapshot = safe_json_dict(getattr(owner, "engine_snapshot_json", None))
chat_snapshot = safe_json_dict(getattr(linked_chat, "engine_snapshot_json", None))
engine = _get_engine_object(record, linked_chat or owner, ctx)
model_config_id = (
getattr(owner, "model_config_id", None)
or getattr(token_usage, "model_config_id", None)
or (record.engine_id if record.engine_type == "model" else None)
or response.get("model_config_id")
)
model_config = ctx.model_configs.get(str(model_config_id)) if model_config_id else None
model_name = (
record.engine_model_name
or response.get("model")
or response.get("model_name")
or chat_snapshot.get("model_name")
or chat_snapshot.get("engine_model_name")
or owner_snapshot.get("model_name")
or owner_snapshot.get("engine_model_name")
or getattr(engine, "model_name", None)
or getattr(model_config, "model_name", None)
)
provider = (
record.engine_provider
or chat_snapshot.get("provider")
or chat_snapshot.get("engine_provider")
or owner_snapshot.get("provider")
or owner_snapshot.get("engine_provider")
or getattr(engine, "provider", None)
or getattr(model_config, "provider", None)
)
if not provider and str(model_name or "").lower().startswith("doubao-"):
provider = "volcengine"
engine_id = (
record.engine_id
or getattr(linked_chat, "engine_id", None)
or getattr(owner, "engine_id", None)
or chat_snapshot.get("engine_id")
or chat_snapshot.get("id")
or owner_snapshot.get("engine_id")
or owner_snapshot.get("id")
or getattr(engine, "id", None)
or getattr(model_config, "id", None)
)
engine_name = (
record.engine_name
or chat_snapshot.get("engine_name")
or chat_snapshot.get("name")
or owner_snapshot.get("engine_name")
or owner_snapshot.get("name")
or getattr(engine, "name", None)
or getattr(model_config, "name", None)
)
record.engine_id = str(engine_id) if engine_id else None
record.engine_name = str(engine_name) if engine_name else None
record.engine_model_name = str(model_name) if model_name else None
record.engine_provider = (
normalize_provider(str(provider), record.engine_model_name)
if provider or record.engine_model_name
else None
)
if not record.engine_type:
if model_config is not None:
record.engine_type = "model"
else:
record.engine_type = str(record.media_type or getattr(linked_chat or owner, "gen_type", None) or "") or None
def _infer_category(record: CreditRecord, owner: Any) -> str | None:
charge_kind = str(record.charge_kind or "").strip()
if charge_kind in {CreditRecordChargeKind.TEXT_PROMPT.value, CreditRecordChargeKind.VIDEO_ANALYSIS.value}:
return "text"
media_type = str(record.media_type or getattr(owner, "gen_type", None) or "").lower()
if media_type in {"image", "video"}:
return media_type
if any(int(value or 0) > 0 for value in (record.input_tokens, record.output_tokens, record.total_tokens)):
return "text"
return None
def _apply_unique_current_rule_fallback(record: CreditRecord, owner: Any, ctx: BackfillContext) -> None:
if record.engine_model_name:
return
category = _infer_category(record, owner)
rules = ctx.current_rules_by_category.get(category or "", [])
unique_models = {(rule.provider, rule.model_name) for rule in rules}
if len(unique_models) != 1:
return
provider, model_name = next(iter(unique_models))
record.engine_provider = provider
record.engine_model_name = model_name
record.engine_type = record.engine_type or ("model" if category == "text" else category)
def _apply_snapshot_fields(
record: CreditRecord,
*,
attachment_snapshot: dict[str, Any] | None,
attachment_counts: dict[str, Any] | None,
generation_snapshot: dict[str, Any] | None,
generation_counts: dict[str, Any] | None,
) -> None:
if attachment_snapshot is not None:
record.attachment_snapshot_json = deepcopy(attachment_snapshot)
for key, value in (attachment_counts or {}).items():
if hasattr(record, key):
setattr(record, key, value)
if generation_snapshot is not None:
record.generation_snapshot_json = deepcopy(generation_snapshot)
for key, value in (generation_counts or {}).items():
if hasattr(record, key):
setattr(record, key, value)
def _merge_resource_counts(
*,
record: CreditRecord,
generation_snapshot: dict[str, Any] | None,
generation_counts: dict[str, Any] | None,
resource_counts: dict[tuple[str, str], dict[str, int]],
) -> tuple[dict[str, Any] | None, dict[str, Any] | None]:
source_model_by_owner_type = {
CreditRecordOwnerType.CHAT_GENERATION_TASK.value: SOURCE_MODEL_CHAT_TASK,
CreditRecordOwnerType.GENERATION_RECORD.value: SOURCE_MODEL_GENERATION_RECORD,
}
source_model = source_model_by_owner_type.get(record.owner_type or "")
bucket = (
resource_counts.get((source_model, record.owner_id))
if source_model and record.owner_id
else None
)
if not bucket:
return generation_snapshot, generation_counts
counts = deepcopy(dict(generation_counts or {}))
counts.update(
{
"generated_image_count": int(bucket["image"]),
"generated_video_count": int(bucket["video"]),
"generated_total_count": int(bucket["image"] + bucket["video"]),
}
)
snapshot = deepcopy(dict(generation_snapshot or {}))
snapshot["generated_image_count"] = counts["generated_image_count"]
snapshot["generated_video_count"] = counts["generated_video_count"]
snapshot["generated_total_count"] = counts["generated_total_count"]
snapshot["output_count_source"] = "generated_resource"
return snapshot, counts
def _ensure_image_output_items(
*,
usage: dict[str, Any],
owner: Any,
successful_count: int,
) -> None:
if successful_count <= 0:
return
items = usage.get("output_items")
if isinstance(items, list) and len(items) >= successful_count:
return
width, height = parse_size(getattr(owner, "image_px", None))
if width <= 0 or height <= 0:
return
usage["output_items"] = [
{
"index": index,
"width": width,
"height": height,
"pixels": width * height,
"size_source": "request_explicit_backfill",
}
for index in range(successful_count)
]
usage["output_pixels_are_estimated"] = True
async def _load_context(
db,
records: list[CreditRecord],
*,
backfill_reference_at: datetime,
) -> BackfillContext:
ctx = BackfillContext()
ids_by_type: dict[str, set[str]] = {}
for record in records:
key = _owner_key(record)
if key:
ids_by_type.setdefault(key[0], set()).add(key[1])
if record.source_step_id:
ids_by_type.setdefault(CreditRecordOwnerType.MODULE_GENERATION_STEP.value, set()).add(record.source_step_id)
model_by_owner = {
CreditRecordOwnerType.CHAT_GENERATION_TASK.value: ChatGenerationTask,
CreditRecordOwnerType.GENERATION_RECORD.value: GenerationRecord,
CreditRecordOwnerType.MODULE_GENERATION_PROJECT.value: ModuleGenerationProject,
CreditRecordOwnerType.MODULE_GENERATION_STEP.value: ModuleGenerationStep,
CreditRecordOwnerType.SHOT_REPLICATE_TASK_SET.value: ShotReplicateTaskSet,
CreditRecordOwnerType.SHOT_REPLICATE_SEGMENT.value: ShotReplicateSegment,
}
for owner_type, model in model_by_owner.items():
ids = ids_by_type.get(owner_type) or set()
if not ids:
ctx.owner_maps[owner_type] = {}
continue
rows = (await db.execute(select(model).where(model.id.in_(ids)))).scalars().all()
ctx.owner_maps[owner_type] = {row.id: row for row in rows}
step_rows = list(ctx.owner_maps.get(CreditRecordOwnerType.MODULE_GENERATION_STEP.value, {}).values())
linked_chat_ids = {str(step.chat_task_id) for step in step_rows if getattr(step, "chat_task_id", None)}
if linked_chat_ids:
rows = (
await db.execute(select(ChatGenerationTask).where(ChatGenerationTask.id.in_(linked_chat_ids)))
).scalars().all()
ctx.linked_chat_tasks = {row.id: row for row in rows}
token_usage_ids = {str(record.token_usage_id) for record in records if record.token_usage_id}
token_usage_ids.update(
str(step.token_usage_id)
for step in step_rows
if getattr(step, "token_usage_id", None)
)
owner_ids = {key[1] for record in records if (key := _owner_key(record))}
token_filters = []
if token_usage_ids:
token_filters.append(TokenUsage.id.in_(token_usage_ids))
if owner_ids:
token_filters.append(TokenUsage.owner_id.in_(owner_ids))
if token_filters:
token_rows = (
await db.execute(
select(TokenUsage)
.where(or_(*token_filters))
.order_by(TokenUsage.created_at.desc())
)
).scalars().all()
for row in token_rows:
ctx.token_usage_by_id[row.id] = row
if row.owner_type and row.owner_id:
ctx.token_usage_by_owner.setdefault((row.owner_type, row.owner_id), row)
model_config_ids = {
str(row.model_config_id)
for row in ctx.token_usage_by_id.values()
if row.model_config_id
}
model_config_ids.update(
str(step.model_config_id)
for step in step_rows
if getattr(step, "model_config_id", None)
)
model_config_ids.update(
str(record.engine_id)
for record in records
if record.engine_type == "model" and record.engine_id
)
if model_config_ids:
rows = (
await db.execute(select(ModelConfig).where(ModelConfig.id.in_(model_config_ids)))
).scalars().all()
ctx.model_configs = {row.id: row for row in rows}
engine_ids = {str(record.engine_id) for record in records if record.engine_id}
for owner_map in ctx.owner_maps.values():
engine_ids.update(str(row.engine_id) for row in owner_map.values() if getattr(row, "engine_id", None))
engine_ids.update(str(row.engine_id) for row in ctx.linked_chat_tasks.values() if row.engine_id)
if engine_ids:
image_rows = (
await db.execute(select(ImageEngine).where(ImageEngine.id.in_(engine_ids)))
).scalars().all()
video_rows = (
await db.execute(select(VideoEngine).where(VideoEngine.id.in_(engine_ids)))
).scalars().all()
ctx.image_engines = {row.id: row for row in image_rows}
ctx.video_engines = {row.id: row for row in video_rows}
source_ids = {
record.owner_id
for record in records
if record.owner_id
and record.owner_type
in {
CreditRecordOwnerType.CHAT_GENERATION_TASK.value,
CreditRecordOwnerType.GENERATION_RECORD.value,
}
}
if source_ids:
resources = (
await db.execute(
select(GeneratedResource)
.where(GeneratedResource.source_id.in_(source_ids))
.where(GeneratedResource.deleted_at.is_(None))
)
).scalars().all()
for resource in resources:
key = (str(resource.source_model or ""), resource.source_id)
bucket = ctx.resource_counts.setdefault(key, {"image": 0, "video": 0})
resource_type = str(resource.resource_type or "").lower()
if resource_type in bucket:
bucket[resource_type] += 1
current_rules = (
await db.execute(
select(ModelPricingRule)
.where(ModelPricingRule.publish_status == ModelPricingRuleStatus.PUBLISHED.value)
.where(ModelPricingRule.effective_from <= backfill_reference_at)
.where(
or_(
ModelPricingRule.effective_to.is_(None),
ModelPricingRule.effective_to > backfill_reference_at,
)
)
.order_by(ModelPricingRule.model_category, ModelPricingRule.model_name)
)
).scalars().all()
for rule in current_rules:
ctx.current_rules_by_category.setdefault(rule.model_category, []).append(rule)
return ctx
def _record_state(record: CreditRecord) -> tuple[Any, ...]:
return (
record.engine_id,
record.engine_provider,
record.engine_model_name,
record.pricing_rule_id,
record.pricing_version_code,
record.provider_cost_status,
record.provider_cost_amount,
record.pricing_snapshot_hash,
record.attachment_total_count,
record.generated_total_count,
deepcopy(record.usage_snapshot_json),
deepcopy(record.attachment_snapshot_json),
deepcopy(record.generation_snapshot_json),
)
def _classify_result(record: CreditRecord, stats: BackfillStats, *, previous_rule_id: str | None, forced: bool) -> None:
status = record.provider_cost_status
if record.pricing_rule_id and record.pricing_rule_id != previous_rule_id:
stats.rule_bound += 1
if status == ProviderCostStatus.CALCULATED.value:
stats.calculated += 1
if forced:
stats.force_repriced += 1
elif status == ProviderCostStatus.ESTIMATED.value:
stats.estimated += 1
if forced:
stats.force_repriced += 1
elif status == ProviderCostStatus.HISTORICAL_ENGINE_UNAVAILABLE.value:
stats.missing_engine += 1
elif status == ProviderCostStatus.USAGE_MISSING.value:
stats.missing_usage += 1
elif status == ProviderCostStatus.UNMATCHED_RULE.value:
stats.unmatched_rule += 1
elif status == ProviderCostStatus.NOT_APPLICABLE.value:
stats.not_applicable += 1
async def run(args) -> None:
start_at = _parse_date(args.start_date)
end_at = _parse_date(args.end_date, end=True)
backfill_reference_at = _utcnow()
total = BackfillStats()
last_id: str | None = None
async with async_session() as db:
while True:
filters = [CreditRecord.type == CreditRecordType.CONSUME.value]
filters.append(
or_(
CreditRecord.charge_action == CreditRecordAction.CHARGE.value,
CreditRecord.charge_action.is_(None),
CreditRecord.charge_action == "",
)
)
filters.append(CreditRecord.created_at <= backfill_reference_at)
if last_id:
filters.append(CreditRecord.id > last_id)
if start_at:
filters.append(CreditRecord.created_at >= start_at)
if end_at:
filters.append(CreditRecord.created_at <= end_at)
if args.only_user_id:
filters.append(CreditRecord.user_id == args.only_user_id)
if args.only_owner_type:
filters.append(CreditRecord.owner_type == args.only_owner_type)
if args.only_missing and not args.force:
filters.append(
or_(
CreditRecord.pricing_rule_id.is_(None),
CreditRecord.provider_cost_status.is_(None),
CreditRecord.provider_cost_status.in_([value for value in INCOMPLETE_COST_STATUSES if value]),
CreditRecord.attachment_snapshot_json.is_(None),
CreditRecord.generation_snapshot_json.is_(None),
)
)
records = (
await db.execute(
select(CreditRecord)
.where(*filters)
.order_by(CreditRecord.id)
.limit(args.batch_size)
.with_for_update(skip_locked=True)
)
).scalars().all()
if not records:
break
ctx = await _load_context(
db,
records,
backfill_reference_at=backfill_reference_at,
)
batch = BackfillStats()
for record in records:
batch.scanned += 1
record_id = str(record.id)
user_id = str(record.user_id) if record.user_id else None
owner_type = str(record.owner_type or "") or None
owner_id = str(record.owner_id or record.related_id or "") or None
if not _is_provider_cost_candidate(record):
batch.skipped_non_provider += 1
continue
if (
not args.force
and record.provider_cost_status == ProviderCostStatus.CALCULATED.value
and record.pricing_rule_id
and record.attachment_snapshot_json is not None
and record.generation_snapshot_json is not None
):
batch.skipped_completed += 1
continue
before = _record_state(record)
previous_rule_id = record.pricing_rule_id
try:
async with db.begin_nested():
key = _owner_key(record)
owner = ctx.owner_maps.get(key[0], {}).get(key[1]) if key else None
if owner is None and record.source_step_id:
owner = ctx.owner_maps.get(
CreditRecordOwnerType.MODULE_GENERATION_STEP.value,
{},
).get(record.source_step_id)
linked_chat = None
if isinstance(owner, ModuleGenerationStep) and owner.chat_task_id:
linked_chat = ctx.linked_chat_tasks.get(owner.chat_task_id)
media_owner = linked_chat or owner
token_usage = _select_token_usage(record, owner, ctx)
usage = _usage_from_record(record, owner, token_usage)
attachment_snapshot, attachment_counts = build_attachment_snapshot(
_raw_references(media_owner or owner)
)
generation_snapshot: dict[str, Any] | None = deepcopy(record.generation_snapshot_json)
generation_counts: dict[str, Any] | None = None
gen_type = str(
record.media_type
or getattr(media_owner, "gen_type", None)
or ""
).lower()
response = _provider_response(media_owner or owner)
if gen_type in {"image", "video"} and media_owner is not None:
generation_snapshot, generation_counts, generation_usage = build_generation_snapshot(
media_owner,
provider_response=response,
stage=PricingSnapshotStage.BACKFILL.value,
)
provider_usage = normalize_provider_media_usage(
response,
gen_type=gen_type,
fallback_total_tokens=max(
safe_int(record.total_tokens),
safe_int(getattr(media_owner, "video_tokens_used", None)),
),
request_image_px=getattr(media_owner, "image_px", None),
requested_output_count=max(
1,
safe_int((generation_counts or {}).get("requested_output_count"), 1),
),
provider_input_image_count=safe_int(
attachment_counts.get("provider_input_image_count")
),
)
usage.update(generation_usage)
usage.update(provider_usage)
usage.update(
{
"has_input_video": bool(
attachment_counts.get("provider_input_video_count")
),
"provider_input_image_count": safe_int(
provider_usage.get("provider_input_image_count"),
safe_int(attachment_counts.get("provider_input_image_count")),
),
"input_image_count": safe_int(
provider_usage.get("provider_input_image_count"),
safe_int(attachment_counts.get("provider_input_image_count")),
),
"input_video_duration_seconds": float(
attachment_counts.get("attachment_video_duration_seconds") or 0
),
"input_audio_duration_seconds": float(
attachment_counts.get("attachment_audio_duration_seconds") or 0
),
"usage_stage": PricingSnapshotStage.BACKFILL.value,
}
)
generation_snapshot, generation_counts = _merge_resource_counts(
record=record,
generation_snapshot=generation_snapshot,
generation_counts=generation_counts,
resource_counts=ctx.resource_counts,
)
if generation_counts:
successful_count = (
int(generation_counts.get("generated_image_count") or 0)
if gen_type == "image"
else int(generation_counts.get("generated_video_count") or 0)
)
usage["successful_output_count"] = successful_count
usage["generated_image_count"] = int(
generation_counts.get("generated_image_count") or 0
)
usage["generated_video_count"] = int(
generation_counts.get("generated_video_count") or 0
)
if gen_type == "image":
_ensure_image_output_items(
usage=usage,
owner=media_owner,
successful_count=successful_count,
)
_restore_engine_snapshot(
record,
owner=owner,
linked_chat=linked_chat,
token_usage=token_usage,
ctx=ctx,
)
_apply_unique_current_rule_fallback(record, media_owner or owner, ctx)
if record.charge_action in {None, ""}:
record.charge_action = CreditRecordAction.CHARGE.value
_apply_snapshot_fields(
record,
attachment_snapshot=attachment_snapshot,
attachment_counts=attachment_counts,
generation_snapshot=generation_snapshot,
generation_counts=generation_counts,
)
await finalize_credit_record_pricing(
db,
charge=record,
usage=usage,
stage=PricingSnapshotStage.BACKFILL.value,
attachment_snapshot=attachment_snapshot,
attachment_counts=attachment_counts,
generation_snapshot=generation_snapshot,
generation_counts=generation_counts,
allow_upgrade_estimated=True,
pricing_reference_at=backfill_reference_at,
use_locked_rule=False,
force_reprice=bool(args.force),
backfill_metadata={
"is_backfilled": True,
"pricing_basis": "current_published_rule",
"backfill_reference_at": backfill_reference_at.isoformat(),
"original_credit_created_at": (
record.created_at.isoformat() if record.created_at else None
),
"command": "backfill_credit_record_snapshots",
},
)
after = _record_state(record)
if before != after:
batch.changed += 1
_classify_result(
record,
batch,
previous_rule_id=previous_rule_id,
forced=bool(args.force),
)
except Exception as exc:
batch.failed += 1
log_model_pricing_event(
event_type="pricing_backfill_record_failed",
event_status="failed",
user_id=user_id,
credit_record_id=record_id,
owner_type=owner_type,
owner_id=owner_id,
error=str(exc),
detail={
"backfill_reference_at": backfill_reference_at.isoformat(),
"force": bool(args.force),
},
)
last_id = str(records[-1].id)
total.merge(batch)
log_model_pricing_event(
event_type="pricing_backfill_batch",
event_status="success" if batch.failed == 0 else "warning",
detail={
**batch.as_dict(),
"batch_size": len(records),
"last_id": last_id,
"commit": bool(args.commit),
"force": bool(args.force),
"backfill_reference_at": backfill_reference_at.isoformat(),
},
)
if args.commit:
await db.commit()
else:
await db.rollback()
print(
f"batch={len(records)} scanned={batch.scanned} changed={batch.changed} "
f"bound={batch.rule_bound} calculated={batch.calculated} estimated={batch.estimated} "
f"missing_engine={batch.missing_engine} missing_usage={batch.missing_usage} "
f"unmatched_rule={batch.unmatched_rule} skipped_completed={batch.skipped_completed} "
f"skipped_non_provider={batch.skipped_non_provider} failed={batch.failed} last_id={last_id}"
)
mode = "COMMIT" if args.commit else "DRY-RUN"
summary = " ".join(f"{key}={value}" for key, value in total.as_dict().items())
print(
f"{mode} DONE backfill_reference_at={backfill_reference_at.isoformat()} "
f"force={bool(args.force)} {summary}"
)
def main() -> None:
parser = argparse.ArgumentParser(
description=(
"一次性将历史消费流水按执行时当前已发布的模型计价规则补齐:"
"同时恢复模型、附件/产出快照、绑定规则并计算供应商成本。"
)
)
group = parser.add_mutually_exclusive_group(required=True)
group.add_argument("--dry-run", action="store_true", help="执行完整计算但最终回滚")
group.add_argument("--commit", action="store_true", help="分批提交补录结果")
parser.add_argument("--batch-size", type=int, default=500)
parser.add_argument("--start-date")
parser.add_argument("--end-date")
parser.add_argument("--only-user-id")
parser.add_argument("--only-owner-type")
parser.add_argument(
"--only-missing",
action="store_true",
default=False,
help="仅扫描规则/成本或附件/产出快照尚未完整的流水",
)
parser.add_argument(
"--force",
action="store_true",
default=False,
help="按当前发布规则覆盖已经核算过的历史计价结果",
)
args = parser.parse_args()
args.batch_size = max(1, min(args.batch_size, 5000))
asyncio.run(run(args))
if __name__ == "__main__":
main()
@@ -0,0 +1,115 @@
from __future__ import annotations
import argparse
import asyncio
from sqlalchemy import select
from app.models.base import async_session
from app.models.menu_config import MenuConfig
from app.models.model_pricing_rule import ModelPricingRule
from app.services.model_pricing.rule_service import create_rule
from app.services.model_pricing.seed_data import volcengine_pricing_seed_rules
from app.utils.id_gen import generate_id
async def _ensure_admin_menu(db) -> bool:
exists = (
await db.execute(
select(MenuConfig.id)
.where(MenuConfig.menu_target == "admin")
.where(MenuConfig.path == "/model-pricing")
.limit(1)
)
).scalar_one_or_none()
if exists:
return False
group_id = (
await db.execute(
select(MenuConfig.id)
.where(MenuConfig.menu_target == "admin")
.where(MenuConfig.menu_type == "group")
.where(MenuConfig.label == "模型设置")
.limit(1)
)
).scalar_one_or_none()
if not group_id:
group_id = generate_id()
db.add(
MenuConfig(
id=group_id,
path="",
label="模型设置",
icon="RobotOutlined",
sort_order=98,
is_active=True,
menu_type="group",
menu_target="admin",
)
)
await db.flush()
db.add(
MenuConfig(
id=generate_id(),
path="/model-pricing",
label="模型计价",
icon="DollarOutlined",
sort_order=4,
is_active=True,
menu_type="page",
menu_target="admin",
parent_id=group_id,
)
)
await db.flush()
return True
async def run(*, commit: bool) -> None:
async with async_session() as db:
created = skipped = 0
menu_created = await _ensure_admin_menu(db)
for payload in volcengine_pricing_seed_rules():
exists = (
await db.execute(
select(ModelPricingRule.id)
.where(ModelPricingRule.provider == payload["provider"])
.where(ModelPricingRule.model_name == payload["model_name"])
.where(ModelPricingRule.version_code == payload["version_code"])
.limit(1)
)
).scalar_one_or_none()
if exists:
skipped += 1
print(f"SKIP {payload['model_name']} {payload['version_code']} id={exists}")
continue
draft_payload = dict(payload)
draft_payload.pop("publish_status", None)
snapshot = await create_rule(db, payload=draft_payload, operator_id=None)
created += 1
effective_from = snapshot["effective_from"]
print(
f"CREATE_DRAFT {snapshot['model_name']} {snapshot['version_code']} id={snapshot['id']} "
f"effective_from={effective_from.isoformat()}"
)
if commit:
await db.commit()
print(f"COMMIT created={created} skipped={skipped} menu_created={menu_created}")
else:
await db.rollback()
print(f"DRY-RUN created={created} skipped={skipped} menu_created={menu_created}")
def main() -> None:
parser = argparse.ArgumentParser(description="初始化火山模型计价草稿(不会自动发布,需人工核价后在后台发布)")
group = parser.add_mutually_exclusive_group(required=True)
group.add_argument("--dry-run", action="store_true")
group.add_argument("--commit", action="store_true")
args = parser.parse_args()
asyncio.run(run(commit=bool(args.commit)))
if __name__ == "__main__":
main()