1
This commit is contained in:
@@ -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()
|
||||
Reference in New Issue
Block a user