from __future__ import annotations import json from datetime import datetime, timedelta, timezone, date from typing import Any from fastapi import HTTPException from sqlalchemy import func, select from sqlalchemy.ext.asyncio import AsyncSession from app.config import settings from app.models.chat_generation_task import ChatGenerationTask from app.models.generation_record import GenerationRecord from app.models.project import Project from app.models.image_engine import ImageEngine from app.models.user import User from app.models.video_engine import VideoEngine from app.enums.generation_history import ( GenerationHistorySourceEnum, get_generation_history_source_label, get_generation_history_task_mode, normalize_generation_history_source, ) from app.schemas.generation_ai import ( GenerationAIEngineGroupOut, GenerationAIEngineOptionsOut, GenerationAIImageEngineOptionOut, GenerationAIRecordHistoryItemOut, GenerationAITaskCreate, GenerationAITaskOut, GenerationAIVideoEngineOptionOut, ) from app.services.generation_billing_service import ( OWNER_CHAT_GENERATION_TASK, charge_generation_media_by_params, ) from app.services.resource_accounting_service import ( SOURCE_MODEL_CHAT_TASK, SOURCE_MODEL_GENERATION_RECORD, batch_get_generated_resource_info_map, soft_delete_chat_task_resources, ) from app.services.resource_signed_url_service import build_resource_signed_url from app.services.generation_history_meta_service import ( GenerationHistoryMeta, batch_load_generation_history_meta_map, build_empty_history_meta, ) from app.services.resource_capacity_service import assert_user_resource_capacity_available from app.utils.id_gen import generate_id IMAGE_DEFAULT_SIZE = "2K" IMAGE_DEFAULT_PROPORTION = "1:1" IMAGE_DEFAULT_PX = "2048x2048" VIDEO_DEFAULT_DURATION = 4 VIDEO_DEFAULT_RATIO = "16:9" VIDEO_DEFAULT_RESOLUTION = "480p" HISTORY_DAY_PAGE_SIZE_MAX = 10 HISTORY_GROUP_ITEM_LIMIT = 10 def normalize_px(value: str | None) -> str | None: if not value: return value return value.replace("×", "x").replace("X", "x").replace("×x", "x").replace("x×", "x") def _json(data: Any) -> str | None: if data is None: return None return json.dumps(data, ensure_ascii=False, default=str) def _parse_json(text: str | None): if not text: return None try: return json.loads(text) except Exception: return None async def _get_image_engine(db: AsyncSession, engine_id: str | None) -> ImageEngine: query = select(ImageEngine).where(ImageEngine.is_active == True) if engine_id: query = query.where(ImageEngine.id == engine_id) else: query = query.order_by(ImageEngine.priority.desc()).limit(1) result = await db.execute(query) engine = result.scalar_one_or_none() if not engine: raise HTTPException(status_code=400, detail="没有可用的图片引擎") return engine async def _get_video_engine(db: AsyncSession, engine_id: str | None) -> VideoEngine: query = select(VideoEngine).where(VideoEngine.is_active == True) if engine_id: query = query.where(VideoEngine.id == engine_id) else: query = query.order_by(VideoEngine.priority.desc()) query = query.limit(1) result = await db.execute(query) engine = result.scalar_one_or_none() if not engine: raise HTTPException(status_code=400, detail="没有可用的视频引擎") return engine def _image_supported_sizes(engine: ImageEngine) -> dict: try: data = json.loads(engine.supported_sizes or "{}") return data if isinstance(data, dict) else {} except Exception: return {} def _parse_list(value: str | None, fallback: list): try: parsed = json.loads(value or "") return parsed if isinstance(parsed, list) else fallback except Exception: return fallback def _build_image_snapshot(engine: ImageEngine, size: str, proportion: str, px: str) -> dict: return { "engine_type": "image", "id": engine.id, "name": engine.name, "provider": engine.provider, "api_base": engine.api_base, "api_key_masked": "****" if engine.api_key else "", "model_name": engine.model_name, "generate_url": engine.generate_url, "supported_models": _parse_list(engine.supported_models, []), "default_size": engine.default_size, "selected_size": size, "selected_proportion": proportion, "selected_px": px, } def _build_video_snapshot(engine: VideoEngine, ratio: str, resolution: str, duration: int) -> dict: return { "engine_type": "video", "id": engine.id, "name": engine.name, "provider": engine.provider, "api_base": engine.api_base, "api_key_masked": "****" if engine.api_key else "", "model_name": engine.model_name, "generate_url": engine.generate_url, "query_url": engine.query_url, "supported_ratios": _parse_list(engine.supported_ratios, []), "supported_resolutions": _parse_list(engine.supported_resolutions, []), "supported_durations": _parse_list(engine.supported_durations, []), "max_duration": engine.max_duration, "selected_ratio": ratio, "selected_resolution": resolution, "selected_duration": duration, } async def list_generation_ai_engine_options(db: AsyncSession) -> GenerationAIEngineOptionsOut: """获取当前启用的图片/视频生成引擎,供前端创建任务时选择 engine_id。""" image_result = await db.execute( select(ImageEngine) .where(ImageEngine.is_active == True) .order_by(ImageEngine.priority.desc()) ) video_result = await db.execute( select(VideoEngine) .where(VideoEngine.is_active == True) .order_by(VideoEngine.priority.desc()) ) image_items = [ GenerationAIImageEngineOptionOut( id=engine.id, name=engine.name, provider=engine.provider, model_name=engine.model_name, supported_models=_parse_list(engine.supported_models, []), supported_sizes=_image_supported_sizes(engine), default_size=engine.default_size, priority=engine.priority or 0, max_image_count=engine.max_image_count, ) for engine in image_result.scalars().all() ] video_items = [ GenerationAIVideoEngineOptionOut( id=engine.id, name=engine.name, provider=engine.provider, model_name=engine.model_name, supported_ratios=_parse_list(engine.supported_ratios, []), supported_resolutions=_parse_list(engine.supported_resolutions, []), supported_durations=_parse_list(engine.supported_durations, []), max_duration=engine.max_duration, priority=engine.priority or 0, max_image_count=engine.max_image_count, max_video_count=engine.max_video_count, supports_first_last_frame=engine.supports_first_last_frame, supports_universal_reference=engine.supports_universal_reference, ) for engine in video_result.scalars().all() ] return GenerationAIEngineOptionsOut( engine=GenerationAIEngineGroupOut(image=image_items, video=video_items) ) async def create_async_generation_task(db: AsyncSession, current_user: User, req: GenerationAITaskCreate) -> ChatGenerationTask: """Create a project-independent chat generation task. Important: this writes chat_generation_tasks, not generation_records, so chat image/video generation no longer needs or validates a project_id. """ gen_type = req.gen_type.lower().strip() if gen_type not in ("image", "video"): raise HTTPException(status_code=400, detail="gen_type 仅支持 image 或 video") if req.idempotency_key: result = await db.execute( select(ChatGenerationTask).where( ChatGenerationTask.user_id == current_user.id, ChatGenerationTask.idempotency_key == req.idempotency_key, ChatGenerationTask.generation_mode == "chatapi_async", ChatGenerationTask.deleted_at.is_(None), ).order_by(ChatGenerationTask.created_at.desc()).limit(1) ) existing = result.scalar_one_or_none() if existing: return existing refs = [r.model_dump() for r in (req.media_references or [])] now = datetime.now(timezone.utc) task_id = generate_id() await assert_user_resource_capacity_available(db, current_user.id) if gen_type == "image": engine = await _get_image_engine(db, req.engine_id) sizes = _image_supported_sizes(engine) size = req.image_size or engine.default_size or IMAGE_DEFAULT_SIZE proportion = req.image_proportion or IMAGE_DEFAULT_PROPORTION px = normalize_px(req.image_px) if sizes: if size not in sizes: raise HTTPException(status_code=400, detail=f"图片分辨率档位不支持: {size}") if proportion not in sizes.get(size, {}): raise HTTPException(status_code=400, detail=f"图片比例不支持: {proportion}") px = px or normalize_px((sizes.get(size) or {}).get(proportion)) px = px or IMAGE_DEFAULT_PX media_billing = await charge_generation_media_by_params( db, user_id=current_user.id, record_id=task_id, gen_type="image", image_size=size, engine_id=engine.id, project_name="AI生成任务", description_prefix="AI创作-", owner_type=OWNER_CHAT_GENERATION_TASK, attempt_no=1, ) snapshot = _build_image_snapshot(engine, size, proportion, px) task = ChatGenerationTask( id=task_id, user_id=current_user.id, original_prompt=req.original_prompt, gen_type="image", image_size=size, image_proportion=proportion, image_px=px, status="generating", generation_mode="chatapi_async", pipeline_stage="queued", engine_id=engine.id, engine_snapshot_json=_json(snapshot), media_references=_json(refs) if refs else None, credits_cost=round(media_billing.total_charged, 2), idempotency_key=req.idempotency_key, deadline_at=now + timedelta(minutes=settings.CHATAPI_ASYNC_IMAGE_DEADLINE_MINUTES), ) else: engine = await _get_video_engine(db, req.engine_id) ratio = req.aspect_ratio or VIDEO_DEFAULT_RATIO resolution = req.resolution or VIDEO_DEFAULT_RESOLUTION duration = req.duration or VIDEO_DEFAULT_DURATION ratios = _parse_list(engine.supported_ratios, []) resolutions = _parse_list(engine.supported_resolutions, []) durations = _parse_list(engine.supported_durations, []) if ratios and ratio not in ratios: raise HTTPException(status_code=400, detail=f"视频比例不支持: {ratio}") if resolutions and resolution not in resolutions: raise HTTPException(status_code=400, detail=f"视频分辨率不支持: {resolution}") if durations and duration not in durations: raise HTTPException(status_code=400, detail=f"视频时长不支持: {duration}") if engine.max_duration and duration > engine.max_duration: raise HTTPException(status_code=400, detail=f"视频时长不能超过 {engine.max_duration} 秒") input_video_duration = 0.0 if refs: video_refs = [r for r in refs if r.get("type") == "video"] for ref in video_refs: ref_duration = float(ref.get("duration") or 0) if ref_duration < 2: raise HTTPException(status_code=400, detail=f"视频素材最短不能少于 2 秒") input_video_duration += ref_duration if input_video_duration > 15: raise HTTPException(status_code=400, detail=f"所有视频素材总时长不能超过 15 秒,当前 {input_video_duration:.1f} 秒") media_billing = await charge_generation_media_by_params( db, user_id=current_user.id, record_id=task_id, gen_type="video", duration=duration, resolution=resolution, engine_id=engine.id, input_video_duration=input_video_duration if input_video_duration > 0 else None, project_name="AI生成任务", description_prefix="AI创作-", owner_type=OWNER_CHAT_GENERATION_TASK, attempt_no=1, ) snapshot = _build_video_snapshot(engine, ratio, resolution, duration) task = ChatGenerationTask( id=task_id, user_id=current_user.id, original_prompt=req.original_prompt, gen_type="video", duration=duration, aspect_ratio=ratio, resolution=resolution, image_size=req.image_size or IMAGE_DEFAULT_SIZE, image_proportion=req.image_proportion or IMAGE_DEFAULT_PROPORTION, image_px=normalize_px(req.image_px) or IMAGE_DEFAULT_PX, status="generating", generation_mode="chatapi_async", pipeline_stage="queued", engine_id=engine.id, engine_snapshot_json=_json(snapshot), media_references=_json(refs) if refs else None, credits_cost=round(media_billing.total_charged, 2), idempotency_key=req.idempotency_key, deadline_at=now + timedelta(hours=settings.CHATAPI_ASYNC_VIDEO_FINAL_DEADLINE_HOURS), ) db.add(task) await db.flush() return task def _resolve_error_message(error_message: str | None) -> str | None: """匹配 ARK_ERRORS 字典,将原始错误码转换为友好提示。 与 app/api/v1/generation.py 的 _record_to_out 保持一致。 注意:celery 任务中已调用 extract_error_message 将错误码转为中文提示后存入数据库, 所以到达此函数的 message 可能是: 1. 已翻译的中文提示(ARK_ERRORS 的 value)→ 直接返回 2. 原始错误字符串(含 code='...' 或 JSON 格式)→ 匹配 ARK_ERRORS 3. 未知内容 → 返回 "生成失败" """ if not error_message: return error_message from app.services.error_codes import ARK_ERRORS # 如果已经是 ARK_ERRORS 中已翻译的中文值,直接返回 if error_message in ARK_ERRORS.values(): return error_message import re # 匹配以下格式中的错误码: # 1. {'error': {'code': 'XXX', ...}} — str(error_obj) 的 Python dict 形式 # 2. {"error": {"code": "XXX", ...}} — JSON 形式 # 3. code='XXX' — 旧格式 for pattern in [ r"'code'\s*:\s*'([^']+)'", # 'code': 'XXX' r'"code"\s*:\s*"([^"]+)"', # "code": "XXX" r"code='([^']+)'", # code='XXX' ]: match = re.search(pattern, error_message) if match: code = match.group(1) if code in ARK_ERRORS: return ARK_ERRORS[code] # 兜底:按冒号分割,检查第二部分是否是已知错误码 parts = error_message.split(":") if len(parts) >= 2 and parts[1].strip() in ARK_ERRORS: return ARK_ERRORS[parts[1].strip()] # 没有匹配到已知错误码时,直接返回"生成失败" return "生成失败" def record_to_out( task: ChatGenerationTask, is_admin: bool = False, generated_resource_id: str | None = None, file_name: str | None = None, history_meta: GenerationHistoryMeta | None = None, ) -> GenerationAITaskOut: refs = _parse_json(task.media_references) snapshot = engine_snapshot_out(_parse_json(task.engine_snapshot_json)) source = GenerationHistorySourceEnum.CHAT_TASK try: source = GenerationHistorySourceEnum( "chat_task" if task.generation_mode == "chatapi_async" else str(task.generation_mode or "chat_task") ) except ValueError: source = GenerationHistorySourceEnum.CHAT_TASK meta = history_meta or build_empty_history_meta(source) return GenerationAITaskOut( id=task.id, user_id=task.user_id if is_admin else None, user_name=getattr(task, "username", None) if is_admin else None, project_id=None, generated_resource_id=generated_resource_id, file_name=file_name, history_source=meta.get("history_source"), history_source_label=meta.get("history_source_label"), module_project_id=meta.get("module_project_id"), module_project_title=meta.get("module_project_title"), module_step_id=meta.get("module_step_id"), module_step_code=meta.get("module_step_code"), hot_opening_project_id=meta.get("hot_opening_project_id"), hot_opening_project_title=meta.get("hot_opening_project_title"), shot_replicate_project_id=meta.get("shot_replicate_project_id"), shot_replicate_project_title=meta.get("shot_replicate_project_title"), shot_task_set_id=meta.get("shot_task_set_id"), shot_segment_id=meta.get("shot_segment_id"), shot_segment_index=meta.get("shot_segment_index"), shot_segment_label=meta.get("shot_segment_label"), gen_type=task.gen_type, generation_mode=task.generation_mode, pipeline_stage=task.pipeline_stage, status=task.status, original_prompt=task.original_prompt, # optimized_prompt=task.optimized_prompt, duration=task.duration, aspect_ratio=task.aspect_ratio, resolution=task.resolution, image_size=task.image_size, image_proportion=task.image_proportion, image_px=task.image_px, media_references=refs, provider_task_id=task.provider_task_id, seedance_task_id=task.seedance_task_id, # remote_result_url=task.remote_result_url, image_url=build_resource_signed_url(task.image_url) if task.image_url else "", video_url=build_resource_signed_url(task.video_url) if task.video_url else "", video_cover_url=build_resource_signed_url(task.video_cover_url) if task.video_cover_url else "", engine_id=task.engine_id, engine_snapshot=snapshot, credits_cost=task.credits_cost or 0.0, text_credits_cost=task.text_credits_cost or 0.0, text_tokens_used=task.text_tokens_used or 0, image_tokens_used=task.image_tokens_used or 0, video_tokens_used=task.video_tokens_used or 0, retry_count=task.retry_count or 0, poll_count=task.poll_count or 0, error_message=_resolve_error_message(task.error_message), created_at=task.created_at, generated_at=task.generated_at, ) def engine_snapshot_out(snapshot: dict) -> dict: """ 从完整的 engine_snapshot 中过滤出需要返回的字段 """ if not snapshot: return {} return { "engine_type": snapshot.get("engine_type"), "id": snapshot.get("id"), "name": snapshot.get("name"), "provider": snapshot.get("provider"), # "api_base": snapshot.get("api_base"), # "api_key_masked": snapshot.get("api_key_masked"), "model_name": snapshot.get("model_name"), # "generate_url": snapshot.get("generate_url"), "supported_models": snapshot.get("supported_models", []), "default_size": snapshot.get("default_size"), "selected_size": snapshot.get("selected_size"), "selected_proportion": snapshot.get("selected_proportion"), "selected_px": snapshot.get("selected_px") } async def list_async_generation_tasks( db: AsyncSession, user_id: str | None, user_name: str | None, gen_type: str | None, status: str | None, page: int, page_size: int, is_admin: bool = False, ): if is_admin: query = ( select(ChatGenerationTask, User.username) .join(User, ChatGenerationTask.user_id == User.id) ) if user_name: query = query.where(User.username.like(f"%{user_name}%")) else: query = select(ChatGenerationTask) query = query.where( ChatGenerationTask.generation_mode == "chatapi_async", ChatGenerationTask.deleted_at.is_(None), ) if user_id: query = query.where(ChatGenerationTask.user_id == user_id) if gen_type: query = query.where(ChatGenerationTask.gen_type == gen_type) if status: query = query.where(ChatGenerationTask.status == status) count_query = select(func.count()).select_from(query.subquery()) total = (await db.execute(count_query)).scalar_one() result = await db.execute( query.order_by(ChatGenerationTask.created_at.desc()) .offset((page - 1) * page_size) .limit(page_size) ) if is_admin: tasks = [] for task, username in result.all(): task.username = username tasks.append(task) return total, tasks return total, list(result.scalars().all()) def _normalize_history_gen_type(gen_type: str | None) -> str: value = (gen_type or "").lower().strip() if value not in ("image", "video"): raise HTTPException(status_code=400, detail="gen_type 仅支持 image 或 video") return value def _normalize_history_source(history_source: str | None) -> GenerationHistorySourceEnum: """Normalize history_source query param. 默认保持原来的 chat_generation_tasks / chatapi_async 历史; 显式传 hot_opening_replicate 或 shot_replicate 时查询对应模块素材; 显式传 generation_record 时查询旧 generation_records 历史。 """ try: return normalize_generation_history_source(history_source) except ValueError: raise HTTPException( status_code=400, detail="history_source 仅支持 chat_task、generation_record、hot_opening_replicate、shot_replicate", ) def _history_day_to_str(value) -> str: if isinstance(value, datetime): return value.date().strftime("%Y-%m-%d") if isinstance(value, date): return value.strftime("%Y-%m-%d") return str(value)[:10] def _parse_history_date(value: str) -> date: try: return datetime.strptime(value, "%Y-%m-%d").date() except ValueError: raise HTTPException(status_code=400, detail="generated_date 格式必须是 YYYY-MM-DD") def _history_base_filters(user_id: str, gen_type: str, source: GenerationHistorySourceEnum): task_mode = get_generation_history_task_mode(source) if not task_mode: raise HTTPException(status_code=400, detail="history_source 不支持查询 ChatGenerationTask 历史") return [ ChatGenerationTask.user_id == user_id, ChatGenerationTask.generation_mode == task_mode.value, ChatGenerationTask.deleted_at.is_(None), ChatGenerationTask.status == "completed", ChatGenerationTask.gen_type == gen_type, ChatGenerationTask.generated_at.is_not(None), ] def _generation_record_history_base_filters(user_id: str, gen_type: str): return [ GenerationRecord.user_id == user_id, GenerationRecord.deleted_at.is_(None), GenerationRecord.status == "completed", GenerationRecord.gen_type == gen_type, GenerationRecord.generated_at.is_not(None), ] def generation_record_to_history_out( record: GenerationRecord, project_name: str | None = None, generated_resource_id: str | None = None, file_name: str | None = None, ) -> GenerationAIRecordHistoryItemOut: refs = _parse_json(record.media_references) return GenerationAIRecordHistoryItemOut( id=record.id, source_type="generation_record", project_id=record.project_id, project_name=project_name, generated_resource_id=generated_resource_id, file_name=file_name, history_source=GenerationHistorySourceEnum.GENERATION_RECORD.value, history_source_label=get_generation_history_source_label(GenerationHistorySourceEnum.GENERATION_RECORD), module_project_id=None, module_project_title=None, module_step_id=None, module_step_code=None, hot_opening_project_id=None, hot_opening_project_title=None, shot_replicate_project_id=None, shot_replicate_project_title=None, shot_task_set_id=None, shot_segment_id=None, shot_segment_index=None, shot_segment_label=None, gen_type=record.gen_type, generation_mode="generation_record", pipeline_stage=None, status=record.status, original_prompt=record.original_prompt, # optimized_prompt=None, duration=record.duration, aspect_ratio=record.aspect_ratio, resolution=record.resolution, image_size=record.image_size, image_proportion=record.image_proportion, image_px=record.image_px, references=refs, media_references=refs, provider_task_id=record.seedance_task_id, seedance_task_id=record.seedance_task_id, remote_result_url=None, image_url=build_resource_signed_url(record.image_url) if record.image_url else '', video_url=build_resource_signed_url(record.video_url) if record.video_url else '', video_cover_url=build_resource_signed_url(record.video_cover_url) if record.video_cover_url else '', engine_id=None, engine_snapshot=None, credits_cost=record.credits_cost or 0.0, text_credits_cost=record.text_credits_cost or 0.0, text_tokens_used=record.text_tokens_used or 0, image_tokens_used=record.image_tokens_used or 0, video_tokens_used=record.video_tokens_used or 0, retry_count=0, poll_count=0, error_message=_resolve_error_message(record.error_message), created_at=record.created_at, generated_at=record.generated_at, ) async def list_generation_record_history_grouped_days( db: AsyncSession, user_id: str, gen_type: str, page: int, page_size: int, ): """ 按生成日期倒序返回旧 generation_records 历史记录分组。 - 每页最多返回 10 个生成日期 - 每个日期分组内最多返回倒序前 10 条旧记录 - 只返回 completed 成功记录 """ gen_type = _normalize_history_gen_type(gen_type) page = max(page, 1) page_size = min(max(page_size, 1), HISTORY_DAY_PAGE_SIZE_MAX) filters = _generation_record_history_base_filters(user_id, gen_type) day_expr = func.date(GenerationRecord.generated_at).label("generated_date") days_subquery = ( select(day_expr) .where(*filters) .group_by(day_expr) .subquery() ) total_days = ( await db.execute(select(func.count()).select_from(days_subquery)) ).scalar_one() day_rows_result = await db.execute( select( day_expr, func.count(GenerationRecord.id).label("total"), ) .where(*filters) .group_by(day_expr) .order_by(day_expr.desc()) .offset((page - 1) * page_size) .limit(page_size) ) day_rows = day_rows_result.all() raw_groups = [] all_record_ids: list[str] = [] for generated_day, day_total in day_rows: item_result = await db.execute( select(GenerationRecord, Project.name.label("project_name")) .outerjoin(Project, (GenerationRecord.project_id == Project.id) & (Project.deleted_at.is_(None))) .where( *filters, func.date(GenerationRecord.generated_at) == generated_day, ) .order_by(GenerationRecord.generated_at.desc(), GenerationRecord.created_at.desc()) .limit(HISTORY_GROUP_ITEM_LIMIT) ) rows = item_result.all() raw_groups.append((generated_day, day_total, rows)) all_record_ids.extend(record.id for record, _project_name in rows) resource_info_map = await batch_get_generated_resource_info_map( db, source_model=SOURCE_MODEL_GENERATION_RECORD, source_ids=all_record_ids, resource_type=gen_type, ) groups = [ { "generated_date": _history_day_to_str(generated_day), "total": int(day_total or 0), "items": [ generation_record_to_history_out( record, project_name, generated_resource_id=resource_info_map.get(record.id, {}).get("resource_id"), file_name=resource_info_map.get(record.id, {}).get("file_name"), ) for record, project_name in rows ], } for generated_day, day_total, rows in raw_groups ] return { "total_days": int(total_days or 0), "page": page, "page_size": page_size, "groups": groups, } async def list_generation_record_history_day_items( db: AsyncSession, user_id: str, gen_type: str, generated_date: str, page: int, page_size: int, ): """ 获取旧 generation_records 指定生成日期下的历史记录分页。 """ gen_type = _normalize_history_gen_type(gen_type) target_day = _parse_history_date(generated_date) page = max(page, 1) page_size = min(max(page_size, 1), 100) filters = _generation_record_history_base_filters(user_id, gen_type) day_expr = func.date(GenerationRecord.generated_at) total = ( await db.execute( select(func.count(GenerationRecord.id)).where( *filters, day_expr == target_day, ) ) ).scalar_one() result = await db.execute( select(GenerationRecord, Project.name.label("project_name")) .outerjoin(Project, (GenerationRecord.project_id == Project.id) & (Project.deleted_at.is_(None))) .where( *filters, day_expr == target_day, ) .order_by(GenerationRecord.generated_at.desc(), GenerationRecord.created_at.desc()) .offset((page - 1) * page_size) .limit(page_size) ) rows = result.all() resource_info_map = await batch_get_generated_resource_info_map( db, source_model=SOURCE_MODEL_GENERATION_RECORD, source_ids=[record.id for record, _project_name in rows], resource_type=gen_type, ) return { "generated_date": target_day.strftime("%Y-%m-%d"), "total": int(total or 0), "page": page, "page_size": page_size, "items": [ generation_record_to_history_out( record, project_name, generated_resource_id=resource_info_map.get(record.id, {}).get("resource_id"), file_name=resource_info_map.get(record.id, {}).get("file_name"), ) for record, project_name in rows ], } async def list_generation_history_grouped_days( db: AsyncSession, user_id: str, gen_type: str, page: int, page_size: int, history_source: str | None = None, ): """ 按生成日期倒序返回历史记录分组。 - 每页最多返回 10 个生成日期 - 每个日期分组内最多返回倒序前 10 条任务 - 只返回 completed 成功任务 - history_source 支持 chat_task / generation_record / hot_opening_replicate / shot_replicate """ source = _normalize_history_source(history_source) if source == GenerationHistorySourceEnum.GENERATION_RECORD: return await list_generation_record_history_grouped_days( db=db, user_id=user_id, gen_type=gen_type, page=page, page_size=page_size, ) gen_type = _normalize_history_gen_type(gen_type) page = max(page, 1) page_size = min(max(page_size, 1), HISTORY_DAY_PAGE_SIZE_MAX) filters = _history_base_filters(user_id, gen_type, source) day_expr = func.date(ChatGenerationTask.generated_at).label("generated_date") days_subquery = ( select(day_expr) .where(*filters) .group_by(day_expr) .subquery() ) total_days = ( await db.execute(select(func.count()).select_from(days_subquery)) ).scalar_one() day_rows_result = await db.execute( select( day_expr, func.count(ChatGenerationTask.id).label("total"), ) .where(*filters) .group_by(day_expr) .order_by(day_expr.desc()) .offset((page - 1) * page_size) .limit(page_size) ) day_rows = day_rows_result.all() raw_groups = [] all_task_ids: list[str] = [] for generated_day, day_total in day_rows: item_result = await db.execute( select(ChatGenerationTask) .where( *filters, func.date(ChatGenerationTask.generated_at) == generated_day, ) .order_by(ChatGenerationTask.generated_at.desc(), ChatGenerationTask.created_at.desc()) .limit(HISTORY_GROUP_ITEM_LIMIT) ) tasks = list(item_result.scalars().all()) raw_groups.append((generated_day, day_total, tasks)) all_task_ids.extend(task.id for task in tasks) resource_info_map = await batch_get_generated_resource_info_map( db, source_model=SOURCE_MODEL_CHAT_TASK, source_ids=all_task_ids, resource_type=gen_type, ) history_meta_map = await batch_load_generation_history_meta_map( db, source=source, chat_task_ids=all_task_ids, ) groups = [ { "generated_date": _history_day_to_str(generated_day), "total": int(day_total or 0), "items": [ record_to_out( task, generated_resource_id=resource_info_map.get(task.id, {}).get("resource_id"), file_name=resource_info_map.get(task.id, {}).get("file_name"), history_meta=history_meta_map.get(task.id), ) for task in tasks ], } for generated_day, day_total, tasks in raw_groups ] return { "total_days": int(total_days or 0), "page": page, "page_size": page_size, "groups": groups, } async def list_generation_history_day_items( db: AsyncSession, user_id: str, gen_type: str, generated_date: str, page: int, page_size: int, history_source: str | None = None, ): """ 获取指定生成日期下的历史记录分页。 用于前端点击某一天后,继续加载该日期下的第 2 页、第 3 页数据。 history_source 支持 chat_task / generation_record / hot_opening_replicate / shot_replicate。 """ source = _normalize_history_source(history_source) if source == GenerationHistorySourceEnum.GENERATION_RECORD: return await list_generation_record_history_day_items( db=db, user_id=user_id, gen_type=gen_type, generated_date=generated_date, page=page, page_size=page_size, ) gen_type = _normalize_history_gen_type(gen_type) target_day = _parse_history_date(generated_date) page = max(page, 1) page_size = min(max(page_size, 1), 100) filters = _history_base_filters(user_id, gen_type, source) day_expr = func.date(ChatGenerationTask.generated_at) total = ( await db.execute( select(func.count(ChatGenerationTask.id)).where( *filters, day_expr == target_day, ) ) ).scalar_one() result = await db.execute( select(ChatGenerationTask) .where( *filters, day_expr == target_day, ) .order_by(ChatGenerationTask.generated_at.desc(), ChatGenerationTask.created_at.desc()) .offset((page - 1) * page_size) .limit(page_size) ) tasks = list(result.scalars().all()) task_ids = [task.id for task in tasks] resource_info_map = await batch_get_generated_resource_info_map( db, source_model=SOURCE_MODEL_CHAT_TASK, source_ids=task_ids, resource_type=gen_type, ) history_meta_map = await batch_load_generation_history_meta_map( db, source=source, chat_task_ids=task_ids, ) return { "generated_date": target_day.strftime("%Y-%m-%d"), "total": int(total or 0), "page": page, "page_size": page_size, "items": [ record_to_out( task, generated_resource_id=resource_info_map.get(task.id, {}).get("resource_id"), file_name=resource_info_map.get(task.id, {}).get("file_name"), history_meta=history_meta_map.get(task.id), ) for task in tasks ], } async def soft_delete_chat_generation_task( db: AsyncSession, *, task: ChatGenerationTask, deleted_at: datetime | None = None, ) -> int: """软删 ChatGenerationTask 并联动软删资源账本,返回释放的 active 空间字节数。""" deleted_at = deleted_at or datetime.now(timezone.utc) task.deleted_at = deleted_at return await soft_delete_chat_task_resources(db, task.id, deleted_at=deleted_at)