Files
video-gen/video-gen-api/app/services/resource_accounting_service.py
T

486 lines
15 KiB
Python

from __future__ import annotations
import json
import os
from dataclasses import dataclass
from datetime import date, datetime, timezone
from typing import Any, Iterable, Sequence
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.chat_generation_task import ChatGenerationTask
from app.models.generated_resource import GeneratedResource
from app.models.generation_record import GenerationRecord
from app.models.user_resource_month_stat import UserResourceMonthStat
from app.models.user_resource_total_stat import UserResourceTotalStat
from app.utils.id_gen import generate_id
SOURCE_MODEL_CHAT_TASK = "ChatGenerationTask"
SOURCE_MODEL_GENERATION_RECORD = "GenerationRecord"
@dataclass(slots=True)
class ResourceAccountingResult:
resource_id: str
resource_type: str
file_size_bytes: int
active_size_delta: int
def resource_month_from_datetime(value: datetime | None = None) -> date:
value = value or datetime.now(timezone.utc)
return date(value.year, value.month, 1)
def safe_file_size(path: str | None) -> int:
if not path:
return 0
try:
return int(os.path.getsize(path))
except OSError:
return 0
def _json(data: Any) -> str | None:
if data is None:
return None
if isinstance(data, str):
return data
return json.dumps(data, ensure_ascii=False, default=str)
def _parse_json(text: str | None) -> dict:
if not text:
return {}
try:
data = json.loads(text)
return data if isinstance(data, dict) else {}
except Exception:
return {}
def _int(value: int | None) -> int:
return int(value or 0)
def _add_non_negative(obj: Any, field: str, delta: int) -> None:
current = _int(getattr(obj, field, 0))
setattr(obj, field, max(0, current + int(delta or 0)))
def _add_raw(obj: Any, field: str, delta: int) -> None:
current = _int(getattr(obj, field, 0))
setattr(obj, field, current + int(delta or 0))
async def _get_or_create_month_stat(
db: AsyncSession,
user_id: str,
stat_month: date,
) -> UserResourceMonthStat:
result = await db.execute(
select(UserResourceMonthStat).where(
UserResourceMonthStat.user_id == user_id,
UserResourceMonthStat.stat_month == stat_month,
)
.limit(1)
)
stat = result.scalar_one_or_none()
if stat:
return stat
stat = UserResourceMonthStat(
id=generate_id(),
user_id=user_id,
stat_month=stat_month,
)
db.add(stat)
await db.flush()
return stat
async def _get_or_create_total_stat(
db: AsyncSession,
user_id: str,
) -> UserResourceTotalStat:
result = await db.execute(
select(UserResourceTotalStat).where(UserResourceTotalStat.user_id == user_id).limit(1)
)
stat = result.scalar_one_or_none()
if stat:
return stat
stat = UserResourceTotalStat(
id=generate_id(),
user_id=user_id,
)
db.add(stat)
await db.flush()
return stat
async def apply_resource_stat_delta(
db: AsyncSession,
*,
user_id: str,
stat_month: date,
resource_type: str,
active_size_delta: int = 0,
active_count_delta: int = 0,
deleted_size_delta: int = 0,
deleted_count_delta: int = 0,
total_generated_size_delta: int = 0,
) -> None:
month_stat = await _get_or_create_month_stat(db, user_id, stat_month)
total_stat = await _get_or_create_total_stat(db, user_id)
now = datetime.now(timezone.utc)
for stat in (month_stat, total_stat):
_add_non_negative(stat, "active_size_bytes", active_size_delta)
_add_non_negative(stat, "active_count", active_count_delta)
_add_non_negative(stat, "deleted_size_bytes", deleted_size_delta)
_add_non_negative(stat, "deleted_count", deleted_count_delta)
_add_raw(stat, "total_generated_size_bytes", total_generated_size_delta)
if resource_type == "image":
_add_non_negative(stat, "image_size_bytes", active_size_delta)
_add_non_negative(stat, "image_count", active_count_delta)
elif resource_type == "video":
_add_non_negative(stat, "video_size_bytes", active_size_delta)
_add_non_negative(stat, "video_count", active_count_delta)
stat.last_recalculated_at = now
async def record_generated_resource(
db: AsyncSession,
*,
user_id: str,
resource_type: str,
resource_url: str,
source_model: str,
source_id: str,
source_model_module: str | None = None,
remote_url: str | None = None,
storage_type: str = "local",
storage_path: str | None = None,
file_size_bytes: int | None = None,
engine_id: str | None = None,
engine_type: str | None = None,
provider: str | None = None,
model_name: str | None = None,
generated_at: datetime | None = None,
extra: Any = None,
) -> ResourceAccountingResult:
"""记录生成成功资源,并增量维护用户月份/全局空间统计。
file_size_bytes 获取不到时允许为 0,符合当前确认方案。
同一个 source_model + source_id + resource_type 已存在未软删账本时,更新账本并按大小差值修正统计。
"""
resource_type = (resource_type or "").lower().strip()
if resource_type not in ("image", "video"):
raise ValueError("resource_type 仅支持 image 或 video")
if not resource_url:
raise ValueError("resource_url 不能为空")
generated_at = generated_at or datetime.now(timezone.utc)
stat_month = resource_month_from_datetime(generated_at)
size = int(file_size_bytes if file_size_bytes is not None else safe_file_size(storage_path))
if size < 0:
size = 0
result = await db.execute(
select(GeneratedResource).where(
GeneratedResource.source_model == source_model,
GeneratedResource.source_id == source_id,
GeneratedResource.resource_type == resource_type,
GeneratedResource.deleted_at.is_(None),
).order_by(GeneratedResource.created_at.desc()).limit(1)
)
existing = result.scalar_one_or_none()
if existing:
old_size = _int(existing.file_size_bytes)
active_size_delta = size - old_size
existing.user_id = user_id
existing.resource_url = resource_url
existing.remote_url = remote_url
existing.storage_type = storage_type or existing.storage_type or "local"
existing.storage_path = storage_path
existing.file_size_bytes = size
existing.engine_id = engine_id
existing.engine_type = engine_type
existing.provider = provider
existing.model_name = model_name
existing.generated_at = generated_at
existing.resource_month = stat_month
existing.extra_json = _json(extra)
if active_size_delta:
await apply_resource_stat_delta(
db,
user_id=user_id,
stat_month=stat_month,
resource_type=resource_type,
active_size_delta=active_size_delta,
)
return ResourceAccountingResult(
resource_id=existing.id,
resource_type=resource_type,
file_size_bytes=size,
active_size_delta=active_size_delta,
)
resource = GeneratedResource(
id=generate_id(),
user_id=user_id,
resource_type=resource_type,
resource_url=resource_url,
remote_url=remote_url,
storage_type=storage_type or "local",
storage_path=storage_path,
file_size_bytes=size,
source_model=source_model,
source_model_module=source_model_module,
source_id=source_id,
engine_id=engine_id,
engine_type=engine_type,
provider=provider,
model_name=model_name,
generated_at=generated_at,
resource_month=stat_month,
extra_json=_json(extra),
)
db.add(resource)
await db.flush()
await apply_resource_stat_delta(
db,
user_id=user_id,
stat_month=stat_month,
resource_type=resource_type,
active_size_delta=size,
active_count_delta=1,
total_generated_size_delta=size,
)
return ResourceAccountingResult(
resource_id=resource.id,
resource_type=resource_type,
file_size_bytes=size,
active_size_delta=size,
)
async def record_chat_task_generated_resource(
db: AsyncSession,
task: ChatGenerationTask,
*,
resource_url: str,
storage_path: str | None = None,
file_size_bytes: int | None = None,
remote_url: str | None = None,
generated_at: datetime | None = None,
) -> ResourceAccountingResult:
snapshot = _parse_json(task.engine_snapshot_json)
return await record_generated_resource(
db,
user_id=task.user_id,
resource_type=task.gen_type,
resource_url=resource_url,
remote_url=remote_url or task.remote_result_url,
storage_type="local" if storage_path else "remote",
storage_path=storage_path,
file_size_bytes=file_size_bytes,
source_model=SOURCE_MODEL_CHAT_TASK,
source_model_module="app.models.chat_generation_task",
source_id=task.id,
engine_id=task.engine_id,
engine_type=snapshot.get("engine_type") or task.gen_type,
provider=snapshot.get("provider"),
model_name=snapshot.get("model_name"),
generated_at=generated_at or task.generated_at or datetime.now(timezone.utc),
extra={"pipeline_stage": task.pipeline_stage},
)
async def record_generation_record_generated_resource(
db: AsyncSession,
record: GenerationRecord,
*,
resource_url: str,
storage_path: str | None = None,
file_size_bytes: int | None = None,
remote_url: str | None = None,
generated_at: datetime | None = None,
) -> ResourceAccountingResult:
return await record_generated_resource(
db,
user_id=record.user_id,
resource_type=record.gen_type,
resource_url=resource_url,
remote_url=remote_url,
storage_type="local" if storage_path else "remote",
storage_path=storage_path,
file_size_bytes=file_size_bytes,
source_model=SOURCE_MODEL_GENERATION_RECORD,
source_model_module="app.models.generation_record",
source_id=record.id,
generated_at=generated_at or record.generated_at or datetime.now(timezone.utc),
extra={"project_id": record.project_id},
)
async def batch_get_generated_resource_id_map(
db: AsyncSession,
*,
source_model: str,
source_ids: Sequence[str] | Iterable[str],
resource_type: str | None = None,
) -> dict[str, str]:
"""批量查询来源记录对应的 GeneratedResource.id。
用于历史列表接口批量回填资源账本 ID,避免按记录一条条查询。
如果历史脏数据存在同一个 source_id 对应多条未软删资源账本,按 created_at 倒序取最新一条。
"""
ids = list(dict.fromkeys(str(item) for item in source_ids if item))
if not ids:
return {}
normalized_resource_type = (resource_type or "").lower().strip()
query = select(GeneratedResource).where(
GeneratedResource.source_model == source_model,
GeneratedResource.source_id.in_(ids),
GeneratedResource.deleted_at.is_(None),
)
if normalized_resource_type in ("image", "video"):
query = query.where(GeneratedResource.resource_type == normalized_resource_type)
result = await db.execute(
query.order_by(
GeneratedResource.source_id.asc(),
GeneratedResource.created_at.desc(),
)
)
resource_id_map: dict[str, str] = {}
for resource in result.scalars().all():
if resource.source_id not in resource_id_map:
resource_id_map[resource.source_id] = resource.id
return resource_id_map
async def batch_get_generated_resource_info_map(
db: AsyncSession,
*,
source_model: str,
source_ids: Sequence[str] | Iterable[str],
resource_type: str | None = None,
) -> dict[str, dict[str, str | None]]:
"""批量查询来源记录对应的 GeneratedResource 信息(id 和 file_name)。
用于历史列表接口批量回填资源账本信息,避免按记录一条条查询。
如果历史脏数据存在同一个 source_id 对应多条未软删资源账本,按 created_at 倒序取最新一条。
返回格式: {source_id: {"resource_id": "...", "file_name": "..."}}
"""
ids = list(dict.fromkeys(str(item) for item in source_ids if item))
if not ids:
return {}
normalized_resource_type = (resource_type or "").lower().strip()
query = select(GeneratedResource).where(
GeneratedResource.source_model == source_model,
GeneratedResource.source_id.in_(ids),
GeneratedResource.deleted_at.is_(None),
)
if normalized_resource_type in ("image", "video"):
query = query.where(GeneratedResource.resource_type == normalized_resource_type)
result = await db.execute(
query.order_by(
GeneratedResource.source_id.asc(),
GeneratedResource.created_at.desc(),
)
)
resource_info_map: dict[str, dict[str, str | None]] = {}
for resource in result.scalars().all():
if resource.source_id not in resource_info_map:
resource_info_map[resource.source_id] = {
"resource_id": resource.id,
"file_name": resource.file_name,
}
return resource_info_map
async def soft_delete_resources_by_source(
db: AsyncSession,
*,
source_model: str,
source_ids: Sequence[str] | Iterable[str],
deleted_at: datetime | None = None,
) -> int:
"""按业务来源软删资源账本,并返回本次释放的 active 空间字节数。"""
ids = [item for item in source_ids if item]
if not ids:
return 0
deleted_at = deleted_at or datetime.now(timezone.utc)
result = await db.execute(
select(GeneratedResource).where(
GeneratedResource.source_model == source_model,
GeneratedResource.source_id.in_(ids),
GeneratedResource.deleted_at.is_(None),
)
)
resources = list(result.scalars().all())
freed_size = 0
for resource in resources:
size = _int(resource.file_size_bytes)
resource.deleted_at = deleted_at
freed_size += size
await apply_resource_stat_delta(
db,
user_id=resource.user_id,
stat_month=resource.resource_month,
resource_type=resource.resource_type,
active_size_delta=-size,
active_count_delta=-1,
deleted_size_delta=size,
deleted_count_delta=1,
)
return freed_size
async def soft_delete_chat_task_resources(
db: AsyncSession,
task_id: str,
*,
deleted_at: datetime | None = None,
) -> int:
return await soft_delete_resources_by_source(
db,
source_model=SOURCE_MODEL_CHAT_TASK,
source_ids=[task_id],
deleted_at=deleted_at,
)
async def soft_delete_generation_record_resources(
db: AsyncSession,
record_ids: Sequence[str] | Iterable[str],
*,
deleted_at: datetime | None = None,
) -> int:
return await soft_delete_resources_by_source(
db,
source_model=SOURCE_MODEL_GENERATION_RECORD,
source_ids=record_ids,
deleted_at=deleted_at,
)