442 lines
14 KiB
Python
442 lines
14 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 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,
|
|
)
|