celery 升级优化V2 | 日志调整 | 前端BUG修复
This commit is contained in:
File diff suppressed because it is too large
Load Diff
@@ -19,3 +19,5 @@ from app.enums.audio_reference import *
|
||||
|
||||
from app.enums.private_portrait import *
|
||||
from app.enums.generation_provider import *
|
||||
|
||||
from app.enums.generation_record import *
|
||||
|
||||
@@ -0,0 +1,31 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from enum import StrEnum
|
||||
|
||||
|
||||
class GenerationRecordConfigSourceEnum(StrEnum):
|
||||
"""生成记录配置冻结来源。"""
|
||||
|
||||
PROMPT_OPTIMIZE = "prompt_optimize"
|
||||
LEGACY_GENERATE_FALLBACK = "legacy_generate_fallback"
|
||||
EXISTING_FROZEN_CONFIG = "existing_frozen_config"
|
||||
|
||||
|
||||
class GenerationRecordEventTypeEnum(StrEnum):
|
||||
"""GenerationRecord 用户生成链路事件。"""
|
||||
|
||||
PROMPT_CONFIG_VALIDATE_START = "PROMPT_CONFIG_VALIDATE_START"
|
||||
PROMPT_CONFIG_VALIDATE_SUCCESS = "PROMPT_CONFIG_VALIDATE_SUCCESS"
|
||||
PROMPT_CONFIG_VALIDATE_FAILED = "PROMPT_CONFIG_VALIDATE_FAILED"
|
||||
PROMPT_CONFIG_FREEZE_START = "PROMPT_CONFIG_FREEZE_START"
|
||||
PROMPT_CONFIG_FREEZE_SUCCESS = "PROMPT_CONFIG_FREEZE_SUCCESS"
|
||||
PROMPT_CONFIG_FREEZE_FAILED = "PROMPT_CONFIG_FREEZE_FAILED"
|
||||
LEGACY_CONFIG_FALLBACK_START = "LEGACY_CONFIG_FALLBACK_START"
|
||||
LEGACY_CONFIG_FALLBACK_SUCCESS = "LEGACY_CONFIG_FALLBACK_SUCCESS"
|
||||
LEGACY_CONFIG_FALLBACK_FAILED = "LEGACY_CONFIG_FALLBACK_FAILED"
|
||||
LEGACY_CONFIG_FALLBACK_SKIPPED = "LEGACY_CONFIG_FALLBACK_SKIPPED"
|
||||
GENERATION_SUBMIT_START = "GENERATION_SUBMIT_START"
|
||||
GENERATION_SUBMIT_CONFIG_READY = "GENERATION_SUBMIT_CONFIG_READY"
|
||||
GENERATION_SUBMIT_BILLING_SUCCESS = "GENERATION_SUBMIT_BILLING_SUCCESS"
|
||||
GENERATION_SUBMIT_ENQUEUE_SUCCESS = "GENERATION_SUBMIT_ENQUEUE_SUCCESS"
|
||||
GENERATION_SUBMIT_FAILED = "GENERATION_SUBMIT_FAILED"
|
||||
@@ -22,6 +22,7 @@ class GenerationRecordPipelineStage(str, Enum):
|
||||
DOWNLOAD_QUEUED = "download_queued"
|
||||
DOWNLOADING = "downloading"
|
||||
RETRY_WAITING = "retry_waiting"
|
||||
RECOVERY_INCONSISTENT = "recovery_inconsistent"
|
||||
UPSCALE_QUEUED = "upscale_queued"
|
||||
UPSCALE_PROCESSING = "upscale_processing"
|
||||
UPSCALE_POLLING = "upscale_polling"
|
||||
|
||||
@@ -47,6 +47,7 @@ class ChatGenerationPipelineStage(str, Enum):
|
||||
DOWNLOAD_QUEUED = "download_queued"
|
||||
DOWNLOADING = "downloading"
|
||||
RETRY_WAITING = "retry_waiting"
|
||||
RECOVERY_INCONSISTENT = "recovery_inconsistent"
|
||||
UPSCALE_QUEUED = "upscale_queued"
|
||||
UPSCALE_PROCESSING = "upscale_processing"
|
||||
UPSCALE_POLLING = "upscale_polling"
|
||||
@@ -107,6 +108,7 @@ class ChatGenerationTaskEventType(str, Enum):
|
||||
FINAL_POLL_BEFORE_TIMEOUT_PENDING = "FINAL_POLL_BEFORE_TIMEOUT_PENDING"
|
||||
GENERATION_RECOVERY_ENQUEUE = "GENERATION_RECOVERY_ENQUEUE"
|
||||
GENERATION_RECOVERY_TIMEOUT = "GENERATION_RECOVERY_TIMEOUT"
|
||||
GENERATION_RECOVERY_INCONSISTENT = "GENERATION_RECOVERY_INCONSISTENT"
|
||||
|
||||
DOWNLOAD_ENQUEUE = "DOWNLOAD_ENQUEUE"
|
||||
DOWNLOAD_ENQUEUE_FAILED = "DOWNLOAD_ENQUEUE_FAILED"
|
||||
|
||||
@@ -1,22 +1,18 @@
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.enums.generation_status import (
|
||||
GenerationStatus,
|
||||
GenerationType,
|
||||
DURATIONS,
|
||||
ASPECT_RATIOS,
|
||||
RESOLUTIONS,
|
||||
IMAGE_SIZES,
|
||||
)
|
||||
from app.enums.generation_status import GenerationType
|
||||
from app.schemas.common import NaiveDatetime, NaiveDatetimeOptional
|
||||
from app.services.operation_log import log_operation
|
||||
|
||||
|
||||
class OptimizeParams(BaseModel):
|
||||
project_id: str
|
||||
prompt: str = Field(..., max_length=500)
|
||||
gen_type: GenerationType = Field(GenerationType.video, description="生成类型:video-视频,image-图片")
|
||||
engine_id: str = Field(..., min_length=1, max_length=32, description="提词阶段选定并冻结的生成引擎ID")
|
||||
include_media_references: bool = Field(False, description="资源生成时是否携带本次提词附件;提词完成后不可修改")
|
||||
duration: int | None = Field(None, description="视频时长(秒),视频生成必填")
|
||||
aspect_ratio: str | None = Field(None, description="视频比例,视频生成必填")
|
||||
resolution: str | None = Field(None, description="视频目标分辨率,视频生成必填")
|
||||
image_size: str | None = Field(None, description="画面分辨率,图片生成使用")
|
||||
image_proportion: str | None = Field(None, description="图片比例,图片生成使用")
|
||||
image_px: str | None = Field(None, description="图片像素大小,图片生成使用")
|
||||
@@ -24,18 +20,10 @@ class OptimizeParams(BaseModel):
|
||||
idempotency_key: str | None = Field(None, max_length=64, description="幂等键,防止重复请求")
|
||||
|
||||
|
||||
class GenerateParams(BaseModel):
|
||||
engine_id: str | None = Field(None, description="生成引擎ID;为空时优先沿用记录引擎,再回退默认引擎")
|
||||
include_media_references: bool = Field(False, description="最终生成时是否携带提词阶段保存的附件")
|
||||
aspect_ratio: str | None = None
|
||||
resolution: str | None = None
|
||||
image_size: str | None = None
|
||||
|
||||
|
||||
class OptimizeResult(BaseModel):
|
||||
optimized_prompt: str
|
||||
text_credits_cost: float
|
||||
# text_tokens_used: int
|
||||
text_tokens_used: int = 0
|
||||
record: "GenerationRecordOut"
|
||||
|
||||
class GenerationRecordOut(BaseModel):
|
||||
@@ -62,11 +50,19 @@ class GenerationRecordOut(BaseModel):
|
||||
engine_name: str | None = None
|
||||
engine_snapshot: dict | None = None
|
||||
include_media_references: bool = False
|
||||
config_complete: bool = False
|
||||
config_recoverable: bool = False
|
||||
config_fallback_hint: str | None = None
|
||||
can_generate: bool = False
|
||||
can_retry: bool = False
|
||||
should_poll: bool = False
|
||||
client_status: str = "ready"
|
||||
operation_phase: str = "prompt"
|
||||
text_credits_cost: float = 0.0
|
||||
# text_tokens_used: int = 0
|
||||
text_tokens_used: int = 0
|
||||
credits_cost: float = 0.0
|
||||
# video_tokens_used: int = 0
|
||||
# image_tokens_used: int = 0
|
||||
video_tokens_used: int = 0
|
||||
image_tokens_used: int = 0
|
||||
error_message: str | None = None
|
||||
created_at: NaiveDatetime
|
||||
generated_at: NaiveDatetimeOptional = None
|
||||
|
||||
@@ -8,9 +8,8 @@ from app.enums.generation_task import GenerationMode, GenerationOwnerType
|
||||
from app.models.base import async_session
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
from app.models.chat_generation_task_event import ChatGenerationTaskEvent
|
||||
from app.models.chat_provider_call_log import ChatProviderCallLog
|
||||
from app.models.generation_record import GenerationRecord
|
||||
from app.services.operation_log_service import build_exception_detail, log_operation_event, sanitize_log_value
|
||||
from app.services.operation_log_service import build_exception_detail, log_ai_model_event, log_operation_event, sanitize_log_value
|
||||
from app.utils.id_gen import generate_id
|
||||
|
||||
MAX_EXCERPT_CHARS = 2000
|
||||
@@ -64,8 +63,14 @@ def _owner_fields(
|
||||
resolved_record_id = None
|
||||
resolved_mode = generation_mode or getattr(obj, "generation_mode", GenerationMode.CHATAPI_ASYNC.value)
|
||||
else:
|
||||
resolved_owner_type = owner_type or GenerationOwnerType.CHAT_GENERATION_TASK.value
|
||||
resolved_owner_id = owner_id
|
||||
inferred_mode = generation_mode or getattr(obj, "generation_mode", None)
|
||||
inferred_owner_type = (
|
||||
GenerationOwnerType.GENERATION_RECORD.value
|
||||
if inferred_mode == GenerationMode.GENERATION_RECORD.value
|
||||
else GenerationOwnerType.CHAT_GENERATION_TASK.value
|
||||
)
|
||||
resolved_owner_type = owner_type or inferred_owner_type
|
||||
resolved_owner_id = owner_id or getattr(obj, "id", None)
|
||||
if resolved_owner_type == GenerationOwnerType.GENERATION_RECORD.value:
|
||||
resolved_task_id = None
|
||||
resolved_record_id = resolved_owner_id
|
||||
@@ -201,8 +206,16 @@ async def log_provider_call(
|
||||
total_tokens: int = 0,
|
||||
error_code: str | None = None,
|
||||
error_message: str | None = None,
|
||||
) -> None:
|
||||
"""Write an owner-scoped provider call log in a separate transaction."""
|
||||
call_id: str | None = None,
|
||||
module: str | None = None,
|
||||
step_code: str | None = None,
|
||||
) -> str | None:
|
||||
"""Write provider audit events to the AiModel file log only.
|
||||
|
||||
``ChatProviderCallLog`` is intentionally no longer written. The model and
|
||||
historical table remain registered for backward compatibility, so no schema
|
||||
migration is required.
|
||||
"""
|
||||
obj = task or record
|
||||
fields = _owner_fields(
|
||||
obj,
|
||||
@@ -214,34 +227,75 @@ async def log_provider_call(
|
||||
generation_mode=generation_mode,
|
||||
)
|
||||
if not fields:
|
||||
return
|
||||
try:
|
||||
async with async_session() as db:
|
||||
db.add(ChatProviderCallLog(
|
||||
id=generate_id(),
|
||||
owner_type=fields["owner_type"],
|
||||
task_id=fields["task_id"],
|
||||
generation_record_id=fields["generation_record_id"],
|
||||
generation_attempt_no=fields["generation_attempt_no"],
|
||||
generation_mode=fields["generation_mode"],
|
||||
provider=provider,
|
||||
api_type=api_type,
|
||||
model=model,
|
||||
engine_id=engine_id,
|
||||
status=status,
|
||||
latency_ms=latency_ms,
|
||||
http_status=http_status,
|
||||
provider_task_id=provider_task_id,
|
||||
request_hash=_hash(request_data),
|
||||
response_hash=_hash(response_data),
|
||||
request_excerpt=_excerpt(request_data),
|
||||
response_excerpt=_excerpt(response_data),
|
||||
prompt_tokens=prompt_tokens or 0,
|
||||
completion_tokens=completion_tokens or 0,
|
||||
total_tokens=total_tokens or 0,
|
||||
error_code=error_code,
|
||||
error_message=error_message,
|
||||
))
|
||||
await db.commit()
|
||||
except Exception as exc:
|
||||
_fallback_log(api_type, fields, exc)
|
||||
return None
|
||||
|
||||
resolved_call_id = call_id or generate_id()
|
||||
resolved_module = module or fields.get("generation_mode") or "generation_pipeline"
|
||||
resolved_step = step_code or api_type
|
||||
common = {
|
||||
"module": resolved_module,
|
||||
"step_code": resolved_step,
|
||||
"call_id": resolved_call_id,
|
||||
"source": "app.services.generation.log_service",
|
||||
"task_id": fields.get("owner_id"),
|
||||
"owner_type": fields.get("owner_type"),
|
||||
"owner_id": fields.get("owner_id"),
|
||||
"generation_attempt_no": fields.get("generation_attempt_no"),
|
||||
"remote_action": api_type,
|
||||
"remote_request_id": provider_task_id,
|
||||
"model_config_id": engine_id,
|
||||
"model_config_name": engine_id,
|
||||
"model_name": model,
|
||||
"provider": provider,
|
||||
"http_status": http_status,
|
||||
}
|
||||
detail = {
|
||||
"generation_mode": fields.get("generation_mode"),
|
||||
"provider_task_id": provider_task_id,
|
||||
"error_code": error_code,
|
||||
}
|
||||
token_usage = {
|
||||
"prompt_tokens": int(prompt_tokens or 0),
|
||||
"completion_tokens": int(completion_tokens or 0),
|
||||
"total_tokens": int(total_tokens or 0),
|
||||
}
|
||||
|
||||
if request_data is not None or str(status).lower() == "request":
|
||||
log_ai_model_event(
|
||||
event_type="REQUEST",
|
||||
event_phase="REQUEST",
|
||||
event_status="started",
|
||||
request=request_data if request_data is not None else {},
|
||||
detail=detail,
|
||||
**common,
|
||||
)
|
||||
|
||||
normalized_status = str(status or "").lower()
|
||||
if response_data is not None or normalized_status in {"success", "completed", "succeeded"}:
|
||||
log_ai_model_event(
|
||||
event_type="RESPONSE",
|
||||
event_phase="RESPONSE",
|
||||
event_status="success" if normalized_status not in {"failed", "error"} else "failed",
|
||||
latency_ms=latency_ms,
|
||||
response=response_data,
|
||||
token_usage=token_usage,
|
||||
detail=detail,
|
||||
error=error_message if normalized_status in {"failed", "error"} else None,
|
||||
**common,
|
||||
)
|
||||
|
||||
if normalized_status in {"failed", "error"} or error_message:
|
||||
log_ai_model_event(
|
||||
event_type="ERROR",
|
||||
event_phase="ERROR",
|
||||
event_status="failed",
|
||||
latency_ms=latency_ms,
|
||||
token_usage=token_usage,
|
||||
detail=build_exception_detail(
|
||||
RuntimeError(error_message or "provider call failed"),
|
||||
detail,
|
||||
),
|
||||
error=error_message or "provider call failed",
|
||||
**common,
|
||||
)
|
||||
return resolved_call_id
|
||||
|
||||
@@ -7,24 +7,50 @@ from app.services.generation.pipeline.owner_service import GenerationOwner, owne
|
||||
|
||||
|
||||
async def enqueue_generation_create(
|
||||
owner: GenerationOwner,
|
||||
owner: GenerationOwner | None = None,
|
||||
*,
|
||||
reason: str,
|
||||
owner_type: str | None = None,
|
||||
owner_id: str | None = None,
|
||||
generation_attempt_no: int | None = None,
|
||||
generation_mode: str | None = None,
|
||||
) -> None:
|
||||
"""Commit caller-owned state before invoking this function."""
|
||||
"""Commit caller-owned state before invoking this function.
|
||||
|
||||
Scalar owner fields are accepted so callers can avoid touching an ORM object after
|
||||
commit. Existing callers may continue passing ``owner``.
|
||||
"""
|
||||
from app.tasks.generation_create_tasks import chatapi_create_generation_task
|
||||
|
||||
owner_type = owner_type_of(owner)
|
||||
attempt_no = int(getattr(owner, "generation_attempt_no", 1) or 1)
|
||||
resolved_owner_type = owner_type or (owner_type_of(owner) if owner is not None else None)
|
||||
resolved_owner_id = owner_id or (str(owner.id) if owner is not None else None)
|
||||
resolved_attempt_no = int(
|
||||
generation_attempt_no
|
||||
or (getattr(owner, "generation_attempt_no", 1) if owner is not None else 1)
|
||||
or 1
|
||||
)
|
||||
if not resolved_owner_type or not resolved_owner_id:
|
||||
raise ValueError("投递生成任务缺少 owner_type 或 owner_id")
|
||||
|
||||
try:
|
||||
chatapi_create_generation_task.apply_async(
|
||||
args=[str(owner.id)],
|
||||
kwargs={"owner_type": owner_type, "generation_attempt_no": attempt_no},
|
||||
args=[resolved_owner_id],
|
||||
kwargs={
|
||||
"owner_type": resolved_owner_type,
|
||||
"generation_attempt_no": resolved_attempt_no,
|
||||
},
|
||||
queue=CeleryQueue.GEN_CHATAPI_CREATE.value,
|
||||
task_id=f"generation-create:{owner_type}:{owner.id}:attempt:{attempt_no}",
|
||||
task_id=(
|
||||
f"generation-create:{resolved_owner_type}:{resolved_owner_id}:"
|
||||
f"attempt:{resolved_attempt_no}"
|
||||
),
|
||||
)
|
||||
await log_task_event(
|
||||
owner,
|
||||
owner_type=resolved_owner_type,
|
||||
owner_id=resolved_owner_id,
|
||||
generation_attempt_no=resolved_attempt_no,
|
||||
generation_mode=generation_mode,
|
||||
event_type=ChatGenerationTaskEventType.GENERATION_RECORD_ENQUEUE_SUCCESS.value,
|
||||
message="资源生成创建任务已投递",
|
||||
detail={"reason": reason, "queue": CeleryQueue.GEN_CHATAPI_CREATE.value},
|
||||
@@ -32,6 +58,10 @@ async def enqueue_generation_create(
|
||||
except Exception as exc:
|
||||
await log_task_event(
|
||||
owner,
|
||||
owner_type=resolved_owner_type,
|
||||
owner_id=resolved_owner_id,
|
||||
generation_attempt_no=resolved_attempt_no,
|
||||
generation_mode=generation_mode,
|
||||
event_type=ChatGenerationTaskEventType.GENERATION_RECORD_ENQUEUE_FAILED.value,
|
||||
message=str(exc),
|
||||
detail={"reason": reason, "queue": CeleryQueue.GEN_CHATAPI_CREATE.value},
|
||||
|
||||
@@ -0,0 +1,491 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from types import SimpleNamespace
|
||||
from typing import Any, Iterable
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.enums.common import LogEventStatusEnum
|
||||
from app.enums.generation_record import (
|
||||
GenerationRecordConfigSourceEnum,
|
||||
GenerationRecordEventTypeEnum,
|
||||
)
|
||||
from app.enums.generation_status import (
|
||||
GenerationType,
|
||||
)
|
||||
from app.models.generation_record import GenerationRecord
|
||||
from app.models.image_engine import ImageEngine
|
||||
from app.models.video_engine import VideoEngine
|
||||
from app.services.generation.ai.engine_service import (
|
||||
IMAGE_DEFAULT_PROPORTION,
|
||||
IMAGE_DEFAULT_PX,
|
||||
IMAGE_DEFAULT_SIZE,
|
||||
VIDEO_DEFAULT_DURATION,
|
||||
VIDEO_DEFAULT_RATIO,
|
||||
VIDEO_DEFAULT_RESOLUTION,
|
||||
image_supported_sizes,
|
||||
normalize_px,
|
||||
parse_json_list,
|
||||
)
|
||||
from app.services.generation.pipeline.generation_record_service import freeze_generation_record_config
|
||||
from app.services.operation_log_service import log_operation_event, log_operation_error
|
||||
from app.services.video_upscale.snapshot_service import build_video_upscale_snapshot
|
||||
from app.utils.exceptions import InvalidStatusError
|
||||
|
||||
|
||||
_GENERATION_RECORD_LOG_DOMAIN = "generation_record"
|
||||
_GENERATION_RECORD_LOG_MODULE = "generation_record"
|
||||
|
||||
|
||||
def _json_loads_object(value: str | None) -> dict[str, Any] | None:
|
||||
if not value:
|
||||
return None
|
||||
try:
|
||||
data = json.loads(value)
|
||||
except (TypeError, json.JSONDecodeError):
|
||||
return None
|
||||
return data if isinstance(data, dict) else None
|
||||
|
||||
|
||||
def _json_loads_list(value: str | None) -> list[Any]:
|
||||
if not value:
|
||||
return []
|
||||
try:
|
||||
data = json.loads(value)
|
||||
except (TypeError, json.JSONDecodeError):
|
||||
return []
|
||||
return data if isinstance(data, list) else []
|
||||
|
||||
|
||||
def generation_record_engine_snapshot(record: GenerationRecord) -> dict[str, Any] | None:
|
||||
return _json_loads_object(record.engine_snapshot_json)
|
||||
|
||||
|
||||
def is_generation_record_config_complete(record: GenerationRecord) -> bool:
|
||||
snapshot = generation_record_engine_snapshot(record)
|
||||
if not record.engine_id or not snapshot:
|
||||
return False
|
||||
if record.gen_type == GenerationType.video.value:
|
||||
return bool(record.duration and record.aspect_ratio and record.resolution)
|
||||
if record.gen_type == GenerationType.image.value:
|
||||
return bool(record.image_size and record.image_proportion and record.image_px)
|
||||
return False
|
||||
|
||||
|
||||
def is_generation_record_config_recoverable(record: GenerationRecord) -> bool:
|
||||
"""Return whether a prompt_optimized legacy row can try server-side config fallback.
|
||||
|
||||
This check intentionally avoids extra DB reads for list pages. The actual engine
|
||||
existence and capability validation is performed while the generate API holds a
|
||||
row lock for the single target record.
|
||||
"""
|
||||
if is_generation_record_config_complete(record):
|
||||
return False
|
||||
if record.status != "prompt_optimized":
|
||||
return False
|
||||
if not record.optimized_prompt:
|
||||
return False
|
||||
return record.gen_type in {GenerationType.video.value, GenerationType.image.value}
|
||||
|
||||
|
||||
def generation_record_config_fallback_hint(record: GenerationRecord) -> str | None:
|
||||
if not is_generation_record_config_recoverable(record):
|
||||
return None
|
||||
return "旧版本记录缺少冻结配置,提交生成时将由后端按可用引擎权重自动补齐一次"
|
||||
|
||||
|
||||
def frozen_generation_record_engine_view(record: GenerationRecord) -> SimpleNamespace:
|
||||
snapshot = generation_record_engine_snapshot(record)
|
||||
if not snapshot:
|
||||
raise InvalidStatusError("该记录缺少冻结的引擎配置,请重新生成提词")
|
||||
snapshot = dict(snapshot)
|
||||
snapshot["id"] = record.engine_id
|
||||
return SimpleNamespace(**snapshot)
|
||||
|
||||
|
||||
def _engine_plain_namespace(engine: ImageEngine | VideoEngine | SimpleNamespace) -> SimpleNamespace:
|
||||
if isinstance(engine, SimpleNamespace):
|
||||
return engine
|
||||
return SimpleNamespace(
|
||||
**{
|
||||
key: value
|
||||
for key, value in vars(engine).items()
|
||||
if key != "_sa_instance_state"
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _safe_int(value: Any) -> int | None:
|
||||
try:
|
||||
return int(value)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
|
||||
def _first_existing_or_default(value: str | None, supported: Iterable[str], default_value: str) -> str:
|
||||
normalized_supported = [str(item).strip() for item in supported if str(item or "").strip()]
|
||||
current = str(value or "").strip()
|
||||
if current and (not normalized_supported or current in normalized_supported):
|
||||
return current
|
||||
if default_value in normalized_supported or not normalized_supported:
|
||||
return default_value
|
||||
return normalized_supported[0]
|
||||
|
||||
|
||||
def _first_duration(value: int | None, supported: Iterable[Any], max_duration: int | None) -> int:
|
||||
supported_ints = [int(item) for item in supported if str(item).isdigit()]
|
||||
current = _safe_int(value)
|
||||
if current and current > 0:
|
||||
if (not supported_ints or current in supported_ints) and (not max_duration or current <= int(max_duration or 0)):
|
||||
return current
|
||||
for item in supported_ints:
|
||||
if item > 0 and (not max_duration or item <= int(max_duration or 0)):
|
||||
return item
|
||||
if max_duration and int(max_duration) > 0:
|
||||
return min(VIDEO_DEFAULT_DURATION, int(max_duration)) or int(max_duration)
|
||||
return VIDEO_DEFAULT_DURATION
|
||||
|
||||
|
||||
def _video_engine_supports_record_params(engine: VideoEngine, record: GenerationRecord) -> bool:
|
||||
ratios = [str(item) for item in parse_json_list(engine.supported_ratios, [])]
|
||||
resolutions = [str(item) for item in parse_json_list(engine.supported_resolutions, [])]
|
||||
durations = [int(item) for item in parse_json_list(engine.supported_durations, []) if str(item).isdigit()]
|
||||
duration = _safe_int(record.duration)
|
||||
if record.aspect_ratio and ratios and record.aspect_ratio not in ratios:
|
||||
return False
|
||||
if record.resolution and resolutions and record.resolution not in resolutions:
|
||||
return False
|
||||
if duration and durations and duration not in durations:
|
||||
return False
|
||||
if duration and int(engine.max_duration or 0) > 0 and duration > int(engine.max_duration or 0):
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _image_engine_supports_record_params(engine: ImageEngine, record: GenerationRecord) -> bool:
|
||||
sizes = image_supported_sizes(engine)
|
||||
if record.image_size and sizes and record.image_size not in sizes:
|
||||
return False
|
||||
if record.image_size and record.image_proportion and sizes:
|
||||
ratios = sizes.get(record.image_size) or {}
|
||||
if ratios and record.image_proportion not in ratios:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
async def _list_active_video_engines(db: AsyncSession) -> list[VideoEngine]:
|
||||
result = await db.execute(
|
||||
select(VideoEngine)
|
||||
.where(VideoEngine.is_active == True, VideoEngine.deleted_at.is_(None))
|
||||
.order_by(VideoEngine.priority.desc(), VideoEngine.created_at.asc(), VideoEngine.id.asc())
|
||||
)
|
||||
return list(result.scalars().all())
|
||||
|
||||
|
||||
async def _list_active_image_engines(db: AsyncSession) -> list[ImageEngine]:
|
||||
result = await db.execute(
|
||||
select(ImageEngine)
|
||||
.where(ImageEngine.is_active == True, ImageEngine.deleted_at.is_(None))
|
||||
.order_by(ImageEngine.priority.desc(), ImageEngine.created_at.asc(), ImageEngine.id.asc())
|
||||
)
|
||||
return list(result.scalars().all())
|
||||
|
||||
|
||||
def _select_video_engine(engines: list[VideoEngine], record: GenerationRecord) -> tuple[VideoEngine, str]:
|
||||
if record.engine_id:
|
||||
for engine in engines:
|
||||
if engine.id == record.engine_id:
|
||||
return engine, "existing_record"
|
||||
for engine in engines:
|
||||
if _video_engine_supports_record_params(engine, record):
|
||||
return engine, "priority_param_match"
|
||||
if engines:
|
||||
return engines[0], "priority_fallback"
|
||||
raise InvalidStatusError("没有可用的视频引擎,无法补齐历史生成配置")
|
||||
|
||||
|
||||
def _select_image_engine(engines: list[ImageEngine], record: GenerationRecord) -> tuple[ImageEngine, str]:
|
||||
if record.engine_id:
|
||||
for engine in engines:
|
||||
if engine.id == record.engine_id:
|
||||
return engine, "existing_record"
|
||||
for engine in engines:
|
||||
if _image_engine_supports_record_params(engine, record):
|
||||
return engine, "priority_param_match"
|
||||
if engines:
|
||||
return engines[0], "priority_fallback"
|
||||
raise InvalidStatusError("没有可用的图片引擎,无法补齐历史生成配置")
|
||||
|
||||
|
||||
def _normalize_video_record_params(record: GenerationRecord, engine: VideoEngine) -> None:
|
||||
ratios = [str(item) for item in parse_json_list(engine.supported_ratios, [])]
|
||||
resolutions = [str(item) for item in parse_json_list(engine.supported_resolutions, [])]
|
||||
durations = parse_json_list(engine.supported_durations, [])
|
||||
record.duration = _first_duration(record.duration, durations, int(engine.max_duration or 0) or None)
|
||||
record.aspect_ratio = _first_existing_or_default(record.aspect_ratio, ratios, VIDEO_DEFAULT_RATIO)
|
||||
record.resolution = _first_existing_or_default(record.resolution, resolutions, VIDEO_DEFAULT_RESOLUTION)
|
||||
|
||||
|
||||
def _normalize_image_record_params(record: GenerationRecord, engine: ImageEngine) -> None:
|
||||
sizes = image_supported_sizes(engine)
|
||||
size_keys = [str(item) for item in sizes.keys() if str(item or "").strip()]
|
||||
current_size = str(record.image_size or "").strip()
|
||||
default_size = str(engine.default_size or IMAGE_DEFAULT_SIZE).strip() or IMAGE_DEFAULT_SIZE
|
||||
if current_size and (not sizes or current_size in sizes):
|
||||
image_size = current_size
|
||||
elif default_size in size_keys:
|
||||
image_size = default_size
|
||||
elif IMAGE_DEFAULT_SIZE in size_keys:
|
||||
image_size = IMAGE_DEFAULT_SIZE
|
||||
elif size_keys:
|
||||
image_size = size_keys[0]
|
||||
else:
|
||||
image_size = current_size or default_size or IMAGE_DEFAULT_SIZE
|
||||
|
||||
ratios = sizes.get(image_size) if sizes else {}
|
||||
ratio_keys = [str(item) for item in (ratios or {}).keys() if str(item or "").strip()]
|
||||
current_ratio = str(record.image_proportion or "").strip()
|
||||
if current_ratio and (not ratio_keys or current_ratio in ratio_keys):
|
||||
image_proportion = current_ratio
|
||||
elif IMAGE_DEFAULT_PROPORTION in ratio_keys or not ratio_keys:
|
||||
image_proportion = IMAGE_DEFAULT_PROPORTION
|
||||
else:
|
||||
image_proportion = ratio_keys[0]
|
||||
|
||||
px_map = ratios or {}
|
||||
current_px = normalize_px(str(record.image_px or "").strip()) if record.image_px else ""
|
||||
image_px = normalize_px(str(px_map.get(image_proportion) or "").strip()) or current_px or IMAGE_DEFAULT_PX
|
||||
|
||||
record.image_size = image_size
|
||||
record.image_proportion = image_proportion
|
||||
record.image_px = image_px
|
||||
|
||||
|
||||
def _config_log_detail(
|
||||
record: GenerationRecord,
|
||||
*,
|
||||
source: str,
|
||||
engine_selected_by: str | None = None,
|
||||
before: dict[str, Any] | None = None,
|
||||
extra: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
references = _json_loads_list(record.media_references)
|
||||
detail: dict[str, Any] = {
|
||||
"record_id": record.id,
|
||||
"user_id": record.user_id,
|
||||
"project_id": record.project_id,
|
||||
"gen_type": record.gen_type,
|
||||
"source": source,
|
||||
"engine_selected_by": engine_selected_by,
|
||||
"engine_id": record.engine_id,
|
||||
"duration": record.duration,
|
||||
"aspect_ratio": record.aspect_ratio,
|
||||
"resolution": record.resolution,
|
||||
"provider_generation_resolution": record.provider_generation_resolution,
|
||||
"image_size": record.image_size,
|
||||
"image_proportion": record.image_proportion,
|
||||
"image_px": record.image_px,
|
||||
"include_media_references": bool(record.include_media_references),
|
||||
"reference_count": len(references),
|
||||
"video_upscale_enabled": bool(record.video_upscale_enabled_snapshot),
|
||||
"config_complete": is_generation_record_config_complete(record),
|
||||
}
|
||||
if before:
|
||||
detail["before"] = before
|
||||
if extra:
|
||||
detail.update(extra)
|
||||
return detail
|
||||
|
||||
|
||||
def _record_config_before(record: GenerationRecord) -> dict[str, Any]:
|
||||
return {
|
||||
"engine_id": record.engine_id,
|
||||
"has_engine_snapshot": bool(record.engine_snapshot_json),
|
||||
"duration": record.duration,
|
||||
"aspect_ratio": record.aspect_ratio,
|
||||
"resolution": record.resolution,
|
||||
"provider_generation_resolution": record.provider_generation_resolution,
|
||||
"image_size": record.image_size,
|
||||
"image_proportion": record.image_proportion,
|
||||
"image_px": record.image_px,
|
||||
"include_media_references": bool(record.include_media_references),
|
||||
}
|
||||
|
||||
|
||||
def log_generation_record_config_event(
|
||||
*,
|
||||
event_type: GenerationRecordEventTypeEnum,
|
||||
event_status: LogEventStatusEnum = LogEventStatusEnum.SUCCESS,
|
||||
source: GenerationRecordConfigSourceEnum | str,
|
||||
record: GenerationRecord,
|
||||
message: str | None = None,
|
||||
detail: dict[str, Any] | None = None,
|
||||
error: str | None = None,
|
||||
) -> None:
|
||||
log_operation_event(
|
||||
domain=_GENERATION_RECORD_LOG_DOMAIN,
|
||||
module=_GENERATION_RECORD_LOG_MODULE,
|
||||
event_type=event_type.value,
|
||||
event_status=event_status.value,
|
||||
source=str(source.value if isinstance(source, GenerationRecordConfigSourceEnum) else source),
|
||||
user_id=str(record.user_id) if record.user_id else None,
|
||||
project_id=str(record.project_id) if record.project_id else None,
|
||||
task_id=str(record.id) if record.id else None,
|
||||
message=message,
|
||||
detail=detail,
|
||||
error=error,
|
||||
)
|
||||
|
||||
|
||||
def freeze_generation_record_config_with_log(
|
||||
record: GenerationRecord,
|
||||
*,
|
||||
engine: ImageEngine | VideoEngine | SimpleNamespace,
|
||||
source: GenerationRecordConfigSourceEnum,
|
||||
) -> None:
|
||||
before = _record_config_before(record)
|
||||
log_generation_record_config_event(
|
||||
event_type=GenerationRecordEventTypeEnum.PROMPT_CONFIG_FREEZE_START,
|
||||
event_status=LogEventStatusEnum.STARTED,
|
||||
source=source,
|
||||
record=record,
|
||||
detail=_config_log_detail(record, source=source.value, before=before),
|
||||
)
|
||||
try:
|
||||
freeze_generation_record_config(record, engine=_engine_plain_namespace(engine))
|
||||
except Exception as exc:
|
||||
log_operation_error(
|
||||
domain=_GENERATION_RECORD_LOG_DOMAIN,
|
||||
event_type=GenerationRecordEventTypeEnum.PROMPT_CONFIG_FREEZE_FAILED.value,
|
||||
module=_GENERATION_RECORD_LOG_MODULE,
|
||||
source=source.value,
|
||||
user_id=str(record.user_id) if record.user_id else None,
|
||||
project_id=str(record.project_id) if record.project_id else None,
|
||||
task_id=str(record.id) if record.id else None,
|
||||
detail=_config_log_detail(record, source=source.value, before=before),
|
||||
exc=exc,
|
||||
)
|
||||
raise
|
||||
log_generation_record_config_event(
|
||||
event_type=GenerationRecordEventTypeEnum.PROMPT_CONFIG_FREEZE_SUCCESS,
|
||||
event_status=LogEventStatusEnum.SUCCESS,
|
||||
source=source,
|
||||
record=record,
|
||||
detail=_config_log_detail(
|
||||
record,
|
||||
source=source.value,
|
||||
before=before,
|
||||
extra={"config_changed": before != _record_config_before(record)},
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
async def ensure_generation_record_config_frozen(
|
||||
db: AsyncSession,
|
||||
record: GenerationRecord,
|
||||
*,
|
||||
source: GenerationRecordConfigSourceEnum = GenerationRecordConfigSourceEnum.LEGACY_GENERATE_FALLBACK,
|
||||
) -> bool:
|
||||
"""Ensure one GenerationRecord has a complete frozen config.
|
||||
|
||||
New records should already be complete and are left untouched. Legacy
|
||||
prompt_optimized rows may be missing engine_id, engine_snapshot_json or
|
||||
selected parameters; those are completed server-side without accepting any
|
||||
generate-time user input.
|
||||
|
||||
Returns True when the record was changed.
|
||||
"""
|
||||
if is_generation_record_config_complete(record):
|
||||
log_generation_record_config_event(
|
||||
event_type=GenerationRecordEventTypeEnum.LEGACY_CONFIG_FALLBACK_SKIPPED,
|
||||
event_status=LogEventStatusEnum.SKIPPED,
|
||||
source=GenerationRecordConfigSourceEnum.EXISTING_FROZEN_CONFIG,
|
||||
record=record,
|
||||
detail=_config_log_detail(record, source=GenerationRecordConfigSourceEnum.EXISTING_FROZEN_CONFIG.value),
|
||||
)
|
||||
return False
|
||||
|
||||
if record.gen_type not in {GenerationType.video.value, GenerationType.image.value}:
|
||||
raise InvalidStatusError("不支持的生成类型,无法补齐历史生成配置")
|
||||
|
||||
before = _record_config_before(record)
|
||||
log_generation_record_config_event(
|
||||
event_type=GenerationRecordEventTypeEnum.LEGACY_CONFIG_FALLBACK_START,
|
||||
event_status=LogEventStatusEnum.STARTED,
|
||||
source=source,
|
||||
record=record,
|
||||
detail=_config_log_detail(record, source=source.value, before=before),
|
||||
)
|
||||
|
||||
try:
|
||||
engine_selected_by = "priority_fallback"
|
||||
if record.gen_type == GenerationType.video.value:
|
||||
engines = await _list_active_video_engines(db)
|
||||
engine, engine_selected_by = _select_video_engine(engines, record)
|
||||
_normalize_video_record_params(record, engine)
|
||||
provider_resolution, upscale_enabled, upscale_snapshot_json = await build_video_upscale_snapshot(
|
||||
db,
|
||||
target_resolution=record.resolution or VIDEO_DEFAULT_RESOLUTION,
|
||||
aspect_ratio=record.aspect_ratio or VIDEO_DEFAULT_RATIO,
|
||||
supported_provider_resolutions=parse_json_list(engine.supported_resolutions, []),
|
||||
)
|
||||
record.provider_generation_resolution = provider_resolution
|
||||
record.video_upscale_enabled_snapshot = upscale_enabled
|
||||
record.video_upscale_snapshot_json = upscale_snapshot_json
|
||||
else:
|
||||
engines = await _list_active_image_engines(db)
|
||||
engine, engine_selected_by = _select_image_engine(engines, record)
|
||||
_normalize_image_record_params(record, engine)
|
||||
record.provider_generation_resolution = None
|
||||
record.video_upscale_enabled_snapshot = False
|
||||
record.video_upscale_snapshot_json = None
|
||||
|
||||
# Historical rows had no explicit resource attachment switch. Missing
|
||||
# values must stay false to avoid silently changing provider input and
|
||||
# billing semantics.
|
||||
record.include_media_references = bool(record.include_media_references)
|
||||
freeze_generation_record_config(record, engine=engine)
|
||||
|
||||
if not is_generation_record_config_complete(record):
|
||||
raise InvalidStatusError("历史生成记录配置自动补齐失败,请重新生成提词")
|
||||
|
||||
log_generation_record_config_event(
|
||||
event_type=GenerationRecordEventTypeEnum.LEGACY_CONFIG_FALLBACK_SUCCESS,
|
||||
event_status=LogEventStatusEnum.SUCCESS,
|
||||
source=source,
|
||||
record=record,
|
||||
detail=_config_log_detail(
|
||||
record,
|
||||
source=source.value,
|
||||
engine_selected_by=engine_selected_by,
|
||||
before=before,
|
||||
extra={"config_changed": before != _record_config_before(record)},
|
||||
),
|
||||
)
|
||||
return before != _record_config_before(record)
|
||||
except HTTPException as exc:
|
||||
log_generation_record_config_event(
|
||||
event_type=GenerationRecordEventTypeEnum.LEGACY_CONFIG_FALLBACK_FAILED,
|
||||
event_status=LogEventStatusEnum.FAILED,
|
||||
source=source,
|
||||
record=record,
|
||||
detail=_config_log_detail(record, source=source.value, before=before),
|
||||
error=str(exc.detail),
|
||||
)
|
||||
raise
|
||||
except Exception as exc:
|
||||
log_operation_error(
|
||||
domain=_GENERATION_RECORD_LOG_DOMAIN,
|
||||
event_type=GenerationRecordEventTypeEnum.LEGACY_CONFIG_FALLBACK_FAILED.value,
|
||||
module=_GENERATION_RECORD_LOG_MODULE,
|
||||
source=source.value,
|
||||
user_id=str(record.user_id) if record.user_id else None,
|
||||
project_id=str(record.project_id) if record.project_id else None,
|
||||
task_id=str(record.id) if record.id else None,
|
||||
detail=_config_log_detail(record, source=source.value, before=before),
|
||||
exc=exc,
|
||||
)
|
||||
raise
|
||||
@@ -18,29 +18,48 @@ def _json(data: dict) -> str:
|
||||
return json.dumps(data, ensure_ascii=False, default=str)
|
||||
|
||||
|
||||
def prepare_generation_record_execution(
|
||||
def freeze_generation_record_config(
|
||||
record: GenerationRecord,
|
||||
*,
|
||||
engine: ImageEngine | VideoEngine,
|
||||
attempt_no: int,
|
||||
) -> None:
|
||||
now = datetime.now(timezone.utc)
|
||||
reset_execution_fields(record, started_at=now, attempt_no=attempt_no)
|
||||
"""Freeze the provider capability and user-selected parameters at prompt time.
|
||||
|
||||
Runtime API keys are intentionally not stored in the snapshot. Provider execution
|
||||
reads only the current secret from the engine row while all capability and selected
|
||||
parameters continue to come from this immutable snapshot.
|
||||
"""
|
||||
record.engine_id = engine.id
|
||||
if record.gen_type == "image":
|
||||
record.engine_snapshot_json = _json(build_image_snapshot(
|
||||
engine,
|
||||
record.image_size or getattr(engine, "default_size", "2K") or "2K",
|
||||
record.image_proportion or "1:1",
|
||||
record.image_px or "2048x2048",
|
||||
))
|
||||
record.engine_snapshot_json = _json(
|
||||
build_image_snapshot(
|
||||
engine,
|
||||
record.image_size or getattr(engine, "default_size", "2K") or "2K",
|
||||
record.image_proportion or "1:1",
|
||||
record.image_px or "2048x2048",
|
||||
)
|
||||
)
|
||||
else:
|
||||
record.engine_snapshot_json = _json(build_video_snapshot(
|
||||
engine,
|
||||
record.aspect_ratio or "16:9",
|
||||
record.resolution or "480p",
|
||||
int(record.duration or 4),
|
||||
))
|
||||
record.engine_snapshot_json = _json(
|
||||
build_video_snapshot(
|
||||
engine,
|
||||
record.aspect_ratio or "16:9",
|
||||
record.resolution or "480p",
|
||||
int(record.duration or 4),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def prepare_generation_record_execution(
|
||||
record: GenerationRecord,
|
||||
*,
|
||||
attempt_no: int,
|
||||
) -> None:
|
||||
"""Reset execution-only fields without changing the frozen prompt configuration."""
|
||||
if not record.engine_id or not record.engine_snapshot_json:
|
||||
raise ValueError("生成记录缺少冻结的引擎配置")
|
||||
now = datetime.now(timezone.utc)
|
||||
reset_execution_fields(record, started_at=now, attempt_no=attempt_no)
|
||||
record.status = GenerationStatus.generating.value
|
||||
record.pipeline_stage = GenerationRecordPipelineStage.QUEUED.value
|
||||
|
||||
@@ -51,9 +70,18 @@ async def commit_and_enqueue_generation_record(
|
||||
*,
|
||||
reason: str,
|
||||
) -> None:
|
||||
record_id = str(record.id)
|
||||
attempt_no = int(record.generation_attempt_no or 1)
|
||||
await db.commit()
|
||||
try:
|
||||
await enqueue_generation_create(record, reason=reason)
|
||||
await enqueue_generation_create(
|
||||
None,
|
||||
reason=reason,
|
||||
owner_type="generation_record",
|
||||
owner_id=record_id,
|
||||
generation_attempt_no=attempt_no,
|
||||
generation_mode="generation_record",
|
||||
)
|
||||
except Exception:
|
||||
# queued stage and all execution metadata are already committed; recovery will retry.
|
||||
# Queued stage and execution metadata are committed; recovery will retry.
|
||||
return
|
||||
|
||||
@@ -24,6 +24,7 @@ class GenerationRecordRecoveryBatch:
|
||||
create: list[GenerationOwnerRef]
|
||||
poll: list[GenerationOwnerRef]
|
||||
download: list[GenerationOwnerRef]
|
||||
inconsistent: list[GenerationOwnerRef]
|
||||
next_cursor: GenerationRecordRecoveryCursor | None
|
||||
|
||||
|
||||
@@ -44,6 +45,7 @@ async def find_generation_record_recovery_batch(
|
||||
GenerationRecordPipelineStage.DOWNLOAD_QUEUED.value,
|
||||
GenerationRecordPipelineStage.DOWNLOADING.value,
|
||||
GenerationRecordPipelineStage.RETRY_WAITING.value,
|
||||
GenerationRecordPipelineStage.RECOVERY_INCONSISTENT.value,
|
||||
}
|
||||
page_size = max(1, int(limit))
|
||||
now = datetime.now(timezone.utc)
|
||||
@@ -77,6 +79,18 @@ async def find_generation_record_recovery_batch(
|
||||
create: list[GenerationOwnerRef] = []
|
||||
poll: list[GenerationOwnerRef] = []
|
||||
download: list[GenerationOwnerRef] = []
|
||||
inconsistent: list[GenerationOwnerRef] = []
|
||||
create_stages = {
|
||||
GenerationRecordPipelineStage.QUEUED.value,
|
||||
GenerationRecordPipelineStage.PREPARING.value,
|
||||
GenerationRecordPipelineStage.CREATING_PROVIDER_TASK.value,
|
||||
}
|
||||
inconsistent_stages = {
|
||||
GenerationRecordPipelineStage.WAITING_REMOTE.value,
|
||||
GenerationRecordPipelineStage.POLLING.value,
|
||||
GenerationRecordPipelineStage.RESULT_READY.value,
|
||||
GenerationRecordPipelineStage.RECOVERY_INCONSISTENT.value,
|
||||
}
|
||||
for (
|
||||
owner_id,
|
||||
attempt_no,
|
||||
@@ -97,7 +111,14 @@ async def find_generation_record_recovery_batch(
|
||||
)
|
||||
stage = str(pipeline_stage or "")
|
||||
if str(remote_result_url or "").strip():
|
||||
if stage == GenerationRecordPipelineStage.RESULT_READY.value:
|
||||
if stage not in {
|
||||
GenerationRecordPipelineStage.DOWNLOAD_QUEUED.value,
|
||||
GenerationRecordPipelineStage.DOWNLOADING.value,
|
||||
GenerationRecordPipelineStage.RETRY_WAITING.value,
|
||||
}:
|
||||
# The remote result URL is stronger recovery evidence than the
|
||||
# persisted stage. Always continue from download instead of
|
||||
# recreating or polling the provider task.
|
||||
download.append(ref)
|
||||
elif stage == GenerationRecordPipelineStage.DOWNLOAD_QUEUED.value:
|
||||
checked_enqueued_at = ensure_aware_utc(download_enqueued_at)
|
||||
@@ -130,7 +151,9 @@ async def find_generation_record_recovery_batch(
|
||||
):
|
||||
poll.append(ref)
|
||||
else:
|
||||
if (
|
||||
if stage in inconsistent_stages:
|
||||
inconsistent.append(ref)
|
||||
elif stage in create_stages and (
|
||||
ensure_aware_utc(provider_create_lease_until) is None
|
||||
or ensure_aware_utc(provider_create_lease_until) <= now
|
||||
):
|
||||
@@ -144,5 +167,6 @@ async def find_generation_record_recovery_batch(
|
||||
create=create,
|
||||
poll=poll,
|
||||
download=download,
|
||||
inconsistent=inconsistent,
|
||||
next_cursor=next_cursor,
|
||||
)
|
||||
|
||||
@@ -1,9 +1,8 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import mimetypes
|
||||
import os
|
||||
import time
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
@@ -91,9 +90,32 @@ async def _get_model_config(db: AsyncSession) -> ModelConfig:
|
||||
|
||||
async def build_prompt_with_chatapi(db: AsyncSession, record: ChatGenerationTask) -> tuple[str, dict]:
|
||||
"""Call ChatAPI once with current request params and attachments. No history context."""
|
||||
config = await _get_model_config(db)
|
||||
if config.provider == "mock":
|
||||
return record.original_prompt, {"input_tokens": 0, "output_tokens": 0, "total_tokens": 0}
|
||||
config_row = await _get_model_config(db)
|
||||
if config_row.provider == "mock":
|
||||
original_prompt = str(record.original_prompt or "")
|
||||
await db.commit()
|
||||
return original_prompt, {"input_tokens": 0, "output_tokens": 0, "total_tokens": 0}
|
||||
|
||||
user_content = await _build_user_content(record, db)
|
||||
config = SimpleNamespace(
|
||||
id=str(config_row.id),
|
||||
name=str(config_row.name or ""),
|
||||
provider=str(config_row.provider or ""),
|
||||
api_base=str(config_row.api_base or ""),
|
||||
api_key=str(config_row.api_key or ""),
|
||||
model_name=str(config_row.model_name or ""),
|
||||
max_tokens=config_row.max_tokens,
|
||||
temperature=config_row.temperature,
|
||||
)
|
||||
record = SimpleNamespace(
|
||||
id=str(record.id),
|
||||
user_id=str(record.user_id),
|
||||
engine_id=str(record.engine_id or "") or None,
|
||||
generation_mode=str(record.generation_mode or ""),
|
||||
generation_attempt_no=int(record.generation_attempt_no or 1),
|
||||
)
|
||||
# Release all configuration/media lookup reads before the remote request.
|
||||
await db.commit()
|
||||
|
||||
system_prompt = (
|
||||
"你是图片/视频生成提示词整理助手。你的职责是根据用户文字、上传图片/视频和生成参数,"
|
||||
@@ -104,14 +126,26 @@ async def build_prompt_with_chatapi(db: AsyncSession, record: ChatGenerationTask
|
||||
"model": config.model_name,
|
||||
"messages": [
|
||||
{"role": "system", "content": system_prompt},
|
||||
{"role": "user", "content": await _build_user_content(record, db)},
|
||||
{"role": "user", "content": user_content},
|
||||
],
|
||||
"max_tokens": config.max_tokens,
|
||||
"temperature": config.temperature,
|
||||
}
|
||||
started = time.perf_counter()
|
||||
call_id = await log_provider_call(
|
||||
record,
|
||||
provider=config.provider,
|
||||
api_type="chat_prompt",
|
||||
model=config.model_name,
|
||||
engine_id=record.engine_id,
|
||||
status="request",
|
||||
request_data=request_data,
|
||||
module="generation_record",
|
||||
step_code="prompt_optimize",
|
||||
)
|
||||
async with provider_limit("ark_chat_prompt", settings.ARK_CHAT_PROMPT_MAX_CONCURRENCY):
|
||||
async with httpx.AsyncClient(timeout=settings.CHATAPI_REQUEST_TIMEOUT_SECONDS) as client:
|
||||
response: httpx.Response | None = None
|
||||
try:
|
||||
response = await client.post(
|
||||
f"{config.api_base.rstrip('/')}/chat/completions",
|
||||
@@ -123,6 +157,7 @@ async def build_prompt_with_chatapi(db: AsyncSession, record: ChatGenerationTask
|
||||
)
|
||||
latency_ms = int((time.perf_counter() - started) * 1000)
|
||||
if response.status_code >= 400:
|
||||
message = response.text[:1000]
|
||||
await log_provider_call(
|
||||
record,
|
||||
provider=config.provider,
|
||||
@@ -132,13 +167,17 @@ async def build_prompt_with_chatapi(db: AsyncSession, record: ChatGenerationTask
|
||||
status="failed",
|
||||
latency_ms=latency_ms,
|
||||
http_status=response.status_code,
|
||||
request_data=request_data,
|
||||
response_data=response.text,
|
||||
error_message=response.text[:1000],
|
||||
error_message=message,
|
||||
call_id=call_id,
|
||||
module="generation_record",
|
||||
step_code="prompt_optimize",
|
||||
)
|
||||
raise RuntimeError(f"ChatAPI HTTP {response.status_code}: {response.text}")
|
||||
raise RuntimeError(f"ChatAPI HTTP {response.status_code}: {message}")
|
||||
data = response.json()
|
||||
except Exception as exc:
|
||||
if isinstance(exc, RuntimeError) and str(exc).startswith("ChatAPI HTTP "):
|
||||
raise
|
||||
latency_ms = int((time.perf_counter() - started) * 1000)
|
||||
await log_provider_call(
|
||||
record,
|
||||
@@ -148,9 +187,12 @@ async def build_prompt_with_chatapi(db: AsyncSession, record: ChatGenerationTask
|
||||
engine_id=record.engine_id,
|
||||
status="failed",
|
||||
latency_ms=latency_ms,
|
||||
request_data=request_data,
|
||||
response_data=None,
|
||||
http_status=response.status_code if response is not None else None,
|
||||
response_data=response.text if response is not None else None,
|
||||
error_message=str(exc),
|
||||
call_id=call_id,
|
||||
module="generation_record",
|
||||
step_code="prompt_optimize",
|
||||
)
|
||||
raise
|
||||
|
||||
@@ -158,23 +200,6 @@ async def build_prompt_with_chatapi(db: AsyncSession, record: ChatGenerationTask
|
||||
input_tokens = int(usage.get("prompt_tokens", 0) or 0)
|
||||
output_tokens = int(usage.get("completion_tokens", 0) or 0)
|
||||
total_tokens = int(usage.get("total_tokens", input_tokens + output_tokens) or 0)
|
||||
content = data.get("choices", [{}])[0].get("message", {}).get("content", "").strip()
|
||||
if not content:
|
||||
raise RuntimeError("ChatAPI未返回有效prompt")
|
||||
|
||||
token_usage_id = generate_id()
|
||||
db.add(TokenUsage(
|
||||
id=token_usage_id,
|
||||
model_config_id=config.id,
|
||||
user_id=record.user_id,
|
||||
owner_type="generation_record",
|
||||
owner_id=record.id,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
total_tokens=total_tokens,
|
||||
))
|
||||
await db.flush()
|
||||
|
||||
await log_provider_call(
|
||||
record,
|
||||
provider=config.provider,
|
||||
@@ -183,13 +208,59 @@ async def build_prompt_with_chatapi(db: AsyncSession, record: ChatGenerationTask
|
||||
engine_id=record.engine_id,
|
||||
status="success",
|
||||
latency_ms=int((time.perf_counter() - started) * 1000),
|
||||
http_status=200,
|
||||
request_data=request_data,
|
||||
http_status=response.status_code if response is not None else 200,
|
||||
response_data=data,
|
||||
prompt_tokens=input_tokens,
|
||||
completion_tokens=output_tokens,
|
||||
total_tokens=total_tokens,
|
||||
call_id=call_id,
|
||||
module="generation_record",
|
||||
step_code="prompt_optimize",
|
||||
)
|
||||
|
||||
content = data.get("choices", [{}])[0].get("message", {}).get("content", "").strip()
|
||||
if not content:
|
||||
await log_provider_call(
|
||||
record,
|
||||
provider=config.provider,
|
||||
api_type="chat_prompt",
|
||||
model=config.model_name,
|
||||
engine_id=record.engine_id,
|
||||
status="failed",
|
||||
error_message="ChatAPI未返回有效prompt",
|
||||
call_id=call_id,
|
||||
module="generation_record",
|
||||
step_code="prompt_optimize",
|
||||
)
|
||||
raise RuntimeError("ChatAPI未返回有效prompt")
|
||||
|
||||
token_usage_id = generate_id()
|
||||
try:
|
||||
db.add(TokenUsage(
|
||||
id=token_usage_id,
|
||||
model_config_id=config.id,
|
||||
user_id=record.user_id,
|
||||
owner_type="generation_record",
|
||||
owner_id=record.id,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
total_tokens=total_tokens,
|
||||
))
|
||||
await db.flush()
|
||||
except Exception as exc:
|
||||
await log_provider_call(
|
||||
record,
|
||||
provider=config.provider,
|
||||
api_type="chat_prompt",
|
||||
model=config.model_name,
|
||||
engine_id=record.engine_id,
|
||||
status="failed",
|
||||
error_message=f"token usage写入失败: {exc}",
|
||||
call_id=call_id,
|
||||
module="generation_record",
|
||||
step_code="prompt_optimize",
|
||||
)
|
||||
raise
|
||||
return content, {
|
||||
"token_usage_id": token_usage_id,
|
||||
"model_config_id": config.id,
|
||||
|
||||
@@ -18,10 +18,11 @@ from app.services.generation.pipeline.owner_service import (
|
||||
from app.models.image_engine import ImageEngine
|
||||
from app.models.video_engine import VideoEngine
|
||||
from app.services.generation.log_service import log_provider_call
|
||||
from app.services.image_gen import ImageProviderError, poll_image_task_status, submit_image_task
|
||||
from app.services.image_gen import poll_image_task_status, submit_image_task
|
||||
from app.services.provider_limit import provider_limit
|
||||
from app.services.video_gen import poll_task_status, submit_video_task
|
||||
from app.types.generation.provider import ImageProviderBatchResult
|
||||
from app.utils.id_gen import generate_id
|
||||
|
||||
|
||||
def _loads(data: str | None) -> dict:
|
||||
@@ -43,6 +44,17 @@ def _try_json(value: Any) -> Any:
|
||||
return None
|
||||
|
||||
|
||||
def _snapshot_owner(task: GenerationOwner) -> SimpleNamespace:
|
||||
"""Copy loaded scalar fields before commit closes the current transaction."""
|
||||
values = {
|
||||
key: value
|
||||
for key, value in vars(task).items()
|
||||
if key != "_sa_instance_state"
|
||||
}
|
||||
values.setdefault("generation_mode", getattr(task, "generation_mode", None) or "generation_record")
|
||||
return SimpleNamespace(**values)
|
||||
|
||||
|
||||
async def get_runtime_engine(db: AsyncSession, task: GenerationOwner) -> Any:
|
||||
"""使用任务快照冻结历史参数,只从当前引擎记录读取密钥。"""
|
||||
snapshot = _loads(task.engine_snapshot_json)
|
||||
@@ -103,37 +115,18 @@ async def create_provider_task(db: AsyncSession, task: GenerationOwner) -> dict:
|
||||
|
||||
async def _create_video_task(db: AsyncSession, task: GenerationOwner) -> dict:
|
||||
engine = await get_runtime_engine(db, task)
|
||||
task_snapshot = _snapshot_owner(task)
|
||||
include_references = owner_include_media_references(task_snapshot)
|
||||
# Close the engine lookup transaction before the long provider HTTP call.
|
||||
await db.commit()
|
||||
started = time.perf_counter()
|
||||
async with provider_limit("ark_video_create", settings.ARK_VIDEO_CREATE_MAX_CONCURRENCY):
|
||||
try:
|
||||
provider_task_id = await submit_video_task(None, engine, task, include_media_references=owner_include_media_references(task))
|
||||
response = {"task_id": provider_task_id}
|
||||
await log_provider_call(
|
||||
task,
|
||||
provider=engine.provider,
|
||||
api_type="video_create",
|
||||
model=engine.model_name,
|
||||
engine_id=task.engine_id,
|
||||
status="success",
|
||||
latency_ms=int((time.perf_counter() - started) * 1000),
|
||||
provider_task_id=provider_task_id,
|
||||
response_data=response,
|
||||
)
|
||||
return {"task_id": provider_task_id, "response_data": response}
|
||||
except Exception as exc:
|
||||
await log_provider_call(
|
||||
task,
|
||||
provider=engine.provider,
|
||||
api_type="video_create",
|
||||
model=engine.model_name,
|
||||
engine_id=task.engine_id,
|
||||
status="failed",
|
||||
latency_ms=int((time.perf_counter() - started) * 1000),
|
||||
error_message=str(exc),
|
||||
)
|
||||
raise
|
||||
provider_task_id = await submit_video_task(
|
||||
None,
|
||||
engine,
|
||||
task_snapshot,
|
||||
include_media_references=include_references,
|
||||
)
|
||||
return {"task_id": provider_task_id, "response_data": {"task_id": provider_task_id}}
|
||||
|
||||
|
||||
async def create_image_sync_batch_result(
|
||||
@@ -143,10 +136,11 @@ async def create_image_sync_batch_result(
|
||||
generation_count: int,
|
||||
) -> ImageProviderBatchResult:
|
||||
engine = await get_runtime_engine(db, task)
|
||||
task_snapshot = _snapshot_owner(task)
|
||||
# Do not keep a database transaction open while the synchronous provider call runs.
|
||||
await db.commit()
|
||||
return await create_image_sync_batch_result_with_engine(
|
||||
task,
|
||||
task_snapshot,
|
||||
engine,
|
||||
generation_count=generation_count,
|
||||
)
|
||||
@@ -163,46 +157,15 @@ async def create_image_sync_batch_result_with_engine(
|
||||
generation_count > 1 时是一次组图 API 调用;失败后绝不退化为多次单图调用。
|
||||
"""
|
||||
count = max(1, int(generation_count or 1))
|
||||
started = time.perf_counter()
|
||||
api_type = "image_sync_batch_create" if count > 1 else "image_sync_create"
|
||||
async with provider_limit("ark_image_sync_create", settings.ARK_IMAGE_CREATE_MAX_CONCURRENCY):
|
||||
try:
|
||||
result = await asyncio.to_thread(
|
||||
submit_image_task,
|
||||
None,
|
||||
engine,
|
||||
task,
|
||||
include_media_references=owner_include_media_references(task),
|
||||
generation_count=count,
|
||||
)
|
||||
response_data = result.get("response_data") or result
|
||||
await log_provider_call(
|
||||
task,
|
||||
provider=engine.provider,
|
||||
api_type=api_type,
|
||||
model=engine.model_name,
|
||||
engine_id=task.engine_id,
|
||||
status="success",
|
||||
latency_ms=int((time.perf_counter() - started) * 1000),
|
||||
provider_task_id=None,
|
||||
response_data=response_data,
|
||||
total_tokens=int(result.get("image_tokens", 0) or 0),
|
||||
)
|
||||
return result
|
||||
except Exception as exc:
|
||||
error_message = exc.safe_message if isinstance(exc, ImageProviderError) else str(exc)
|
||||
await log_provider_call(
|
||||
task,
|
||||
provider=engine.provider,
|
||||
api_type=api_type,
|
||||
model=engine.model_name,
|
||||
engine_id=task.engine_id,
|
||||
status="failed",
|
||||
latency_ms=int((time.perf_counter() - started) * 1000),
|
||||
error_message=error_message,
|
||||
response_data=exc.as_dict() if isinstance(exc, ImageProviderError) else None,
|
||||
)
|
||||
raise
|
||||
return await asyncio.to_thread(
|
||||
submit_image_task,
|
||||
None,
|
||||
engine,
|
||||
task,
|
||||
include_media_references=owner_include_media_references(task),
|
||||
generation_count=count,
|
||||
)
|
||||
|
||||
|
||||
async def create_image_sync_result(db: AsyncSession, task: GenerationOwner) -> dict:
|
||||
@@ -226,13 +189,59 @@ async def create_image_sync_result(db: AsyncSession, task: GenerationOwner) -> d
|
||||
|
||||
async def poll_provider_task(db: AsyncSession, task: GenerationOwner) -> dict:
|
||||
engine = await get_runtime_engine(db, task)
|
||||
task_snapshot = _snapshot_owner(task)
|
||||
task_id = owner_provider_task_id(task_snapshot)
|
||||
# Polling may block on the remote provider; release the lookup transaction first.
|
||||
await db.commit()
|
||||
task_id = owner_provider_task_id(task)
|
||||
if not task_id:
|
||||
raise ValueError("缺少供应商任务ID")
|
||||
if task.gen_type == "video":
|
||||
async with provider_limit("ark_video_poll", settings.ARK_VIDEO_POLL_MAX_CONCURRENCY):
|
||||
return await poll_task_status(engine, task_id)
|
||||
async with provider_limit("ark_image_poll", settings.ARK_IMAGE_POLL_MAX_CONCURRENCY):
|
||||
return await poll_image_task_status(engine, task_id)
|
||||
|
||||
api_type = f"{task_snapshot.gen_type}_poll"
|
||||
call_id = generate_id()
|
||||
await log_provider_call(
|
||||
task_snapshot,
|
||||
provider=engine.provider,
|
||||
api_type=api_type,
|
||||
model=engine.model_name,
|
||||
engine_id=task_snapshot.engine_id,
|
||||
status="request",
|
||||
provider_task_id=task_id,
|
||||
request_data={"provider_task_id": task_id},
|
||||
call_id=call_id,
|
||||
)
|
||||
started = time.perf_counter()
|
||||
try:
|
||||
if task_snapshot.gen_type == "video":
|
||||
async with provider_limit("ark_video_poll", settings.ARK_VIDEO_POLL_MAX_CONCURRENCY):
|
||||
result = await poll_task_status(engine, task_id)
|
||||
else:
|
||||
async with provider_limit("ark_image_poll", settings.ARK_IMAGE_POLL_MAX_CONCURRENCY):
|
||||
result = await poll_image_task_status(engine, task_id)
|
||||
await log_provider_call(
|
||||
task_snapshot,
|
||||
provider=engine.provider,
|
||||
api_type=api_type,
|
||||
model=engine.model_name,
|
||||
engine_id=task_snapshot.engine_id,
|
||||
status="success",
|
||||
latency_ms=int((time.perf_counter() - started) * 1000),
|
||||
provider_task_id=task_id,
|
||||
response_data=_try_json(result.get("response_data")) or result,
|
||||
total_tokens=int(result.get("video_tokens", 0) or result.get("image_tokens", 0) or 0),
|
||||
call_id=call_id,
|
||||
)
|
||||
return result
|
||||
except Exception as exc:
|
||||
await log_provider_call(
|
||||
task_snapshot,
|
||||
provider=engine.provider,
|
||||
api_type=api_type,
|
||||
model=engine.model_name,
|
||||
engine_id=task_snapshot.engine_id,
|
||||
status="failed",
|
||||
latency_ms=int((time.perf_counter() - started) * 1000),
|
||||
provider_task_id=task_id,
|
||||
error_message=str(exc),
|
||||
call_id=call_id,
|
||||
)
|
||||
raise
|
||||
|
||||
@@ -95,6 +95,45 @@ async def _load_chat_task_for_update(
|
||||
return owner if isinstance(owner, ChatGenerationTask) else None
|
||||
|
||||
|
||||
def _chat_task_post_commit_snapshot(task: ChatGenerationTask) -> Any:
|
||||
"""Capture fields used by Redis/Celery/logging before committing the ORM row."""
|
||||
from types import SimpleNamespace
|
||||
|
||||
return SimpleNamespace(
|
||||
id=str(task.id),
|
||||
generation_attempt_no=int(task.generation_attempt_no or 1),
|
||||
generation_mode=str(task.generation_mode or GenerationMode.CHATAPI_ASYNC.value),
|
||||
provider_task_id=str(task.provider_task_id or "") or None,
|
||||
seedance_task_id=str(task.seedance_task_id or "") or None,
|
||||
gen_type=str(task.gen_type or ""),
|
||||
pipeline_stage=str(task.pipeline_stage or ""),
|
||||
poll_count=int(task.poll_count or 0),
|
||||
poll_error_count=int(task.poll_error_count or 0),
|
||||
manual_retry_count=int(task.manual_retry_count or 0),
|
||||
poll_started_at=task.poll_started_at,
|
||||
poll_interval_seconds=int(task.poll_interval_seconds or 0),
|
||||
last_poll_at=task.last_poll_at,
|
||||
next_poll_at=task.next_poll_at,
|
||||
poll_lease_until=task.poll_lease_until,
|
||||
deadline_at=task.deadline_at,
|
||||
user_id=str(task.user_id or "") or None,
|
||||
project_id=str(task.project_id or "") or None,
|
||||
error_message=str(task.error_message or "") or None,
|
||||
)
|
||||
|
||||
|
||||
async def _reload_chat_task_after_commit(
|
||||
db: AsyncSession, task_id: str
|
||||
) -> ChatGenerationTask | None:
|
||||
owner = await load_generation_owner(
|
||||
db,
|
||||
owner_type=GenerationOwnerType.CHAT_GENERATION_TASK.value,
|
||||
owner_id=str(task_id),
|
||||
for_update=False,
|
||||
)
|
||||
return owner if isinstance(owner, ChatGenerationTask) else None
|
||||
|
||||
|
||||
def _now() -> datetime:
|
||||
return datetime.now(timezone.utc)
|
||||
|
||||
@@ -125,6 +164,7 @@ def _is_final_task_state(task: ChatGenerationTask) -> bool:
|
||||
ChatGenerationPipelineStage.FAILED.value,
|
||||
ChatGenerationPipelineStage.TIMEOUT.value,
|
||||
ChatGenerationPipelineStage.DOWNLOAD_FAILED.value,
|
||||
ChatGenerationPipelineStage.UPSCALE_FAILED.value,
|
||||
)
|
||||
|
||||
|
||||
@@ -430,11 +470,17 @@ async def _mark_timeout(
|
||||
error_message=error_message,
|
||||
pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value,
|
||||
)
|
||||
snapshot = _chat_task_post_commit_snapshot(task)
|
||||
await db.commit()
|
||||
await notify_owner_finished(db, task)
|
||||
await _remove_poll_active(_chat_registry_id(task))
|
||||
fresh_task = await _reload_chat_task_after_commit(db, snapshot.id)
|
||||
if fresh_task is not None:
|
||||
await notify_owner_finished(db, fresh_task)
|
||||
await _remove_poll_active(_chat_registry_id(snapshot))
|
||||
await log_task_event(
|
||||
task,
|
||||
owner_type=GenerationOwnerType.CHAT_GENERATION_TASK.value,
|
||||
owner_id=snapshot.id,
|
||||
generation_attempt_no=snapshot.generation_attempt_no,
|
||||
generation_mode=snapshot.generation_mode,
|
||||
event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value,
|
||||
to_status="failed",
|
||||
to_stage=ChatGenerationPipelineStage.TIMEOUT.value,
|
||||
@@ -456,10 +502,21 @@ async def _mark_failed(
|
||||
error_message=error_message,
|
||||
pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
|
||||
)
|
||||
snapshot = _chat_task_post_commit_snapshot(task)
|
||||
await db.commit()
|
||||
await notify_owner_finished(db, task)
|
||||
await _remove_poll_active(_chat_registry_id(task))
|
||||
await log_task_event(task, event_type=event_type, message=task.error_message, detail=detail)
|
||||
fresh_task = await _reload_chat_task_after_commit(db, snapshot.id)
|
||||
if fresh_task is not None:
|
||||
await notify_owner_finished(db, fresh_task)
|
||||
await _remove_poll_active(_chat_registry_id(snapshot))
|
||||
await log_task_event(
|
||||
owner_type=GenerationOwnerType.CHAT_GENERATION_TASK.value,
|
||||
owner_id=snapshot.id,
|
||||
generation_attempt_no=snapshot.generation_attempt_no,
|
||||
generation_mode=snapshot.generation_mode,
|
||||
event_type=event_type,
|
||||
message=snapshot.error_message or error_message,
|
||||
detail=detail,
|
||||
)
|
||||
return "mark_failed"
|
||||
|
||||
|
||||
@@ -475,8 +532,8 @@ async def recover_one_generation_task(
|
||||
分流原则:
|
||||
1. 已有 remote_result_url:只恢复下载,不 poll,不重新 create。
|
||||
2. 已有 provider_task_id/seedance_task_id:恢复 poll。
|
||||
3. 无结果 URL、无供应商任务 ID:deadline 未过才恢复 create。
|
||||
4. 无结果 URL、无供应商任务 ID:deadline 已过直接超时失败,不再补救生成。
|
||||
3. 仅 queued/preparing/creating_provider_task 且无远程证据时允许恢复 create。
|
||||
4. waiting_remote/polling/result_ready 缺少对应证据时隔离,deadline 到期后失败退款。
|
||||
"""
|
||||
from app.tasks.generation_create_tasks import chatapi_create_generation_task
|
||||
from app.tasks.generation_download_tasks import enqueue_download_task
|
||||
@@ -543,21 +600,25 @@ async def recover_one_generation_task(
|
||||
if is_deadline_expired:
|
||||
if has_provider_task_id:
|
||||
task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value
|
||||
snapshot = _chat_task_post_commit_snapshot(task)
|
||||
await db.commit()
|
||||
await log_task_event(
|
||||
task,
|
||||
owner_type=GenerationOwnerType.CHAT_GENERATION_TASK.value,
|
||||
owner_id=snapshot.id,
|
||||
generation_attempt_no=snapshot.generation_attempt_no,
|
||||
generation_mode=snapshot.generation_mode,
|
||||
event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value,
|
||||
message=f"{source} 发现任务已到 deadline 且存在供应商任务ID,投递 poll 队列做最终查询",
|
||||
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
|
||||
detail={"pipeline_stage": snapshot.pipeline_stage, "payload": redis_payload},
|
||||
)
|
||||
poll_generation_task.apply_async(
|
||||
args=[task.id],
|
||||
kwargs={"force_due": True, "owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": int(task.generation_attempt_no or 1)},
|
||||
args=[snapshot.id],
|
||||
kwargs={"force_due": True, "owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": snapshot.generation_attempt_no},
|
||||
queue=POLL_QUEUE,
|
||||
countdown=0,
|
||||
)
|
||||
await register_poll_active(
|
||||
task,
|
||||
snapshot,
|
||||
check_at=_poll_queue_timeout_at(),
|
||||
reason=f"{source}_deadline_final_poll",
|
||||
)
|
||||
@@ -575,21 +636,25 @@ async def recover_one_generation_task(
|
||||
if is_video_generation_task(task):
|
||||
ensure_video_poll_fields(task, now=current_time)
|
||||
if is_poll_not_due(task, now=current_time):
|
||||
snapshot = _chat_task_post_commit_snapshot(task)
|
||||
await db.commit()
|
||||
await register_poll_active(
|
||||
task,
|
||||
check_at=task.next_poll_at,
|
||||
next_poll_at=task.next_poll_at,
|
||||
snapshot,
|
||||
check_at=snapshot.next_poll_at,
|
||||
next_poll_at=snapshot.next_poll_at,
|
||||
reason=f"{source}_video_poll_not_due",
|
||||
)
|
||||
await log_task_event(
|
||||
task,
|
||||
owner_type=GenerationOwnerType.CHAT_GENERATION_TASK.value,
|
||||
owner_id=snapshot.id,
|
||||
generation_attempt_no=snapshot.generation_attempt_no,
|
||||
generation_mode=snapshot.generation_mode,
|
||||
event_type=ChatGenerationTaskEventType.POLL_SKIP_NOT_DUE.value,
|
||||
message=f"{source} 发现视频任务尚未到下一次轮询时间,启动容灾不提前投递 poll",
|
||||
detail={
|
||||
"pipeline_stage": task.pipeline_stage,
|
||||
"pipeline_stage": snapshot.pipeline_stage,
|
||||
"payload": redis_payload,
|
||||
"next_poll_at": task.next_poll_at,
|
||||
"next_poll_at": snapshot.next_poll_at,
|
||||
},
|
||||
)
|
||||
return "skip_video_poll_not_due"
|
||||
@@ -600,28 +665,32 @@ async def recover_one_generation_task(
|
||||
# 这里仍复用 next_poll_at 做短暂队列保护,避免启动容灾重复投递。
|
||||
# 真正消费时通过 force_due=True 跳过“未到期”校验,避免保护时间反向阻塞本次 poll。
|
||||
task.next_poll_at = queue_hold_until
|
||||
snapshot = _chat_task_post_commit_snapshot(task)
|
||||
await db.commit()
|
||||
await log_task_event(
|
||||
task,
|
||||
owner_type=GenerationOwnerType.CHAT_GENERATION_TASK.value,
|
||||
owner_id=snapshot.id,
|
||||
generation_attempt_no=snapshot.generation_attempt_no,
|
||||
generation_mode=snapshot.generation_mode,
|
||||
event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value,
|
||||
message=f"{source} 发现任务存在供应商任务ID,恢复投递轮询队列",
|
||||
detail={
|
||||
"pipeline_stage": task.pipeline_stage,
|
||||
"pipeline_stage": snapshot.pipeline_stage,
|
||||
"payload": redis_payload,
|
||||
"due_next_poll_at": original_next_poll_at,
|
||||
"queue_hold_until": queue_hold_until,
|
||||
},
|
||||
)
|
||||
poll_generation_task.apply_async(
|
||||
args=[task.id],
|
||||
kwargs={"force_due": True, "owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": int(task.generation_attempt_no or 1)},
|
||||
args=[snapshot.id],
|
||||
kwargs={"force_due": True, "owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": snapshot.generation_attempt_no},
|
||||
queue=POLL_QUEUE,
|
||||
countdown=0,
|
||||
)
|
||||
await register_poll_active(
|
||||
task,
|
||||
check_at=task.next_poll_at,
|
||||
next_poll_at=task.next_poll_at,
|
||||
snapshot,
|
||||
check_at=snapshot.next_poll_at,
|
||||
next_poll_at=snapshot.next_poll_at,
|
||||
reason=f"{source}_has_provider_task_id",
|
||||
)
|
||||
return "recover_poll_has_provider_id"
|
||||
@@ -633,64 +702,84 @@ async def recover_one_generation_task(
|
||||
ChatGenerationPipelineStage.QUEUED.value,
|
||||
ChatGenerationPipelineStage.PREPARING.value,
|
||||
ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
|
||||
ChatGenerationPipelineStage.WAITING_REMOTE.value,
|
||||
ChatGenerationPipelineStage.POLLING.value,
|
||||
}
|
||||
if task.pipeline_stage in recoverable_create_stages:
|
||||
if task.pipeline_stage not in (
|
||||
ChatGenerationPipelineStage.QUEUED.value,
|
||||
ChatGenerationPipelineStage.PREPARING.value,
|
||||
ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
|
||||
):
|
||||
task.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value
|
||||
# 刷新更新时间形成创建队列保护窗口,避免 Beat 在任务尚未消费时每轮重复补投。
|
||||
task.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value
|
||||
task.updated_at = current_time
|
||||
# Release the recovery row lock before writing an event through the
|
||||
# independent logging session or talking to the broker.
|
||||
task_id = str(task.id)
|
||||
attempt_no = int(task.generation_attempt_no or 1)
|
||||
generation_mode = str(task.generation_mode or "")
|
||||
await db.commit()
|
||||
|
||||
await _remove_poll_active(_chat_registry_id(task))
|
||||
await _remove_poll_active(
|
||||
redis_owner_item_id(
|
||||
GenerationOwnerType.CHAT_GENERATION_TASK.value,
|
||||
task_id,
|
||||
attempt_no,
|
||||
)
|
||||
)
|
||||
await log_task_event(
|
||||
task,
|
||||
owner_type=GenerationOwnerType.CHAT_GENERATION_TASK.value,
|
||||
owner_id=task_id,
|
||||
generation_attempt_no=attempt_no,
|
||||
generation_mode=generation_mode,
|
||||
event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value,
|
||||
message=f"{source} 发现任务未超时且缺少 remote_result_url/供应商任务ID,恢复投递创建队列",
|
||||
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
|
||||
detail={
|
||||
"pipeline_stage": ChatGenerationPipelineStage.QUEUED.value,
|
||||
"payload": redis_payload,
|
||||
},
|
||||
)
|
||||
chatapi_create_generation_task.apply_async(
|
||||
args=[task.id],
|
||||
kwargs={"owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": int(task.generation_attempt_no or 1)},
|
||||
args=[task_id],
|
||||
kwargs={
|
||||
"owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value,
|
||||
"generation_attempt_no": attempt_no,
|
||||
},
|
||||
queue=CeleryQueue.GEN_CHATAPI_CREATE.value,
|
||||
countdown=0,
|
||||
task_id=(
|
||||
f"generation-create:{GenerationOwnerType.CHAT_GENERATION_TASK.value}:"
|
||||
f"{task.id}:attempt:{int(task.generation_attempt_no or 1)}"
|
||||
f"{task_id}:attempt:{attempt_no}"
|
||||
),
|
||||
)
|
||||
return "recover_create_no_remote_no_provider_before_deadline"
|
||||
|
||||
# result_ready 但没有 URL 是脏状态;未过 deadline 时回创建队列重新处理,过期上面已标记超时。
|
||||
if task.pipeline_stage == ChatGenerationPipelineStage.RESULT_READY.value:
|
||||
task.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value
|
||||
task.updated_at = current_time
|
||||
inconsistent_stages = {
|
||||
ChatGenerationPipelineStage.WAITING_REMOTE.value,
|
||||
ChatGenerationPipelineStage.POLLING.value,
|
||||
ChatGenerationPipelineStage.RESULT_READY.value,
|
||||
ChatGenerationPipelineStage.RECOVERY_INCONSISTENT.value,
|
||||
}
|
||||
if task.pipeline_stage in inconsistent_stages:
|
||||
original_stage = str(task.pipeline_stage or "")
|
||||
task_id = str(task.id)
|
||||
attempt_no = int(task.generation_attempt_no or 1)
|
||||
generation_mode = str(task.generation_mode or "")
|
||||
task.pipeline_stage = ChatGenerationPipelineStage.RECOVERY_INCONSISTENT.value
|
||||
task.error_message = (
|
||||
f"{source} 恢复证据异常:阶段 {original_stage} 缺少 remote_result_url 和供应商任务ID"
|
||||
)
|
||||
await db.commit()
|
||||
await _remove_poll_active(_chat_registry_id(task))
|
||||
await _remove_poll_active(
|
||||
redis_owner_item_id(
|
||||
GenerationOwnerType.CHAT_GENERATION_TASK.value,
|
||||
task_id,
|
||||
attempt_no,
|
||||
)
|
||||
)
|
||||
await log_task_event(
|
||||
task,
|
||||
event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value,
|
||||
message=f"{source} 发现 result_ready 但缺少 remote_result_url,未超时,恢复投递创建队列",
|
||||
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
|
||||
owner_type=GenerationOwnerType.CHAT_GENERATION_TASK.value,
|
||||
owner_id=task_id,
|
||||
generation_attempt_no=attempt_no,
|
||||
generation_mode=generation_mode,
|
||||
event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_INCONSISTENT.value,
|
||||
from_stage=original_stage,
|
||||
to_stage=ChatGenerationPipelineStage.RECOVERY_INCONSISTENT.value,
|
||||
message=f"{source} 发现恢复证据异常,已隔离且不重新创建供应商任务",
|
||||
detail={"payload": redis_payload},
|
||||
)
|
||||
chatapi_create_generation_task.apply_async(
|
||||
args=[task.id],
|
||||
kwargs={"owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": int(task.generation_attempt_no or 1)},
|
||||
queue=CeleryQueue.GEN_CHATAPI_CREATE.value,
|
||||
countdown=0,
|
||||
task_id=(
|
||||
f"generation-create:{GenerationOwnerType.CHAT_GENERATION_TASK.value}:"
|
||||
f"{task.id}:attempt:{int(task.generation_attempt_no or 1)}"
|
||||
),
|
||||
)
|
||||
return "recover_create_result_ready_no_url_before_deadline"
|
||||
return "quarantine_inconsistent_recovery_evidence"
|
||||
|
||||
return f"skip_stage_{task.pipeline_stage}"
|
||||
|
||||
@@ -919,6 +1008,7 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
|
||||
"waiting_remote",
|
||||
"polling",
|
||||
"result_ready",
|
||||
"recovery_inconsistent",
|
||||
]
|
||||
),
|
||||
)
|
||||
|
||||
@@ -971,6 +971,13 @@ async def run_image_prompt_optimize(
|
||||
user_id=user_id_value,
|
||||
references=references,
|
||||
gen_type="image",
|
||||
log_module=module_value,
|
||||
log_step="hot_opening_image_prompt_optimize",
|
||||
log_project_id=project_id_value,
|
||||
log_task_id=step_id_value,
|
||||
log_owner_type="module_generation_step",
|
||||
log_owner_id=step_id_value,
|
||||
generation_attempt_no=expected_step_version,
|
||||
)
|
||||
if execution_guard is not None:
|
||||
await execution_guard()
|
||||
|
||||
@@ -3,6 +3,7 @@ from __future__ import annotations
|
||||
import copy
|
||||
import json
|
||||
import re
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
@@ -15,7 +16,6 @@ from app.enums.hot_opening_replicate import HotOpeningLogEventEnum, HotOpeningRe
|
||||
from app.enums.shot_replicate import ModuleCodeEnum as ShotModuleCodeEnum, ShotReplicateLogEventEnum, ShotReplicateRemoteActionEnum
|
||||
from app.services.operation_log_service import log_ai_model_event
|
||||
from app.enums.common import (
|
||||
VIDEO_SCHEMA_CONFIG_DATABASE_SOURCE,
|
||||
VIDEO_SCHEMA_CONFIG_DEFAULT_SOURCE,
|
||||
VIDEO_SCHEMA_CONFIG_VERSION,
|
||||
VIDEO_SCHEMA_EDITABLE_TEXT_MAX_LEN,
|
||||
@@ -1477,6 +1477,7 @@ def _log_video_prompt_ai_event(
|
||||
event_status: str,
|
||||
config: ModelConfig,
|
||||
trace_id: str | None,
|
||||
call_id: str,
|
||||
user_id: str | None,
|
||||
project_id: str | None,
|
||||
step_id: str | None,
|
||||
@@ -1489,30 +1490,61 @@ def _log_video_prompt_ai_event(
|
||||
error: str | None = None,
|
||||
detail: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
log_ai_model_event(
|
||||
event_type=event_type,
|
||||
event_status=event_status,
|
||||
source=LogSourceEnum.REMOTE_API.value,
|
||||
module=module,
|
||||
trace_id=trace_id,
|
||||
user_id=user_id,
|
||||
project_id=project_id,
|
||||
step_id=step_id,
|
||||
remote_action=action,
|
||||
remote_request_id=remote_request_id,
|
||||
model_config_id=str(config.id),
|
||||
model_config_name=config.name,
|
||||
model_name=config.model_name,
|
||||
provider=config.provider,
|
||||
api_base=config.api_base,
|
||||
http_status=http_status,
|
||||
request=request_data,
|
||||
response=response_data,
|
||||
token_usage=token_usage,
|
||||
message=message,
|
||||
detail=detail,
|
||||
error=error,
|
||||
)
|
||||
common = {
|
||||
"source": LogSourceEnum.REMOTE_API.value,
|
||||
"module": module,
|
||||
"step_code": "video_prompt_generate",
|
||||
"call_id": call_id,
|
||||
"trace_id": trace_id,
|
||||
"user_id": user_id,
|
||||
"project_id": project_id,
|
||||
"task_id": step_id,
|
||||
"step_id": step_id,
|
||||
"owner_type": "module_generation_step",
|
||||
"owner_id": step_id or project_id,
|
||||
"remote_action": action,
|
||||
"remote_request_id": remote_request_id,
|
||||
"model_config_id": str(config.id),
|
||||
"model_config_name": config.name,
|
||||
"model_name": config.model_name,
|
||||
"provider": config.provider,
|
||||
"api_base": config.api_base,
|
||||
"http_status": http_status,
|
||||
}
|
||||
normalized_status = str(event_status or "").lower()
|
||||
if normalized_status == str(LogEventStatusEnum.STARTED.value).lower():
|
||||
log_ai_model_event(
|
||||
event_type=event_type,
|
||||
event_phase="REQUEST",
|
||||
event_status=event_status,
|
||||
request=request_data,
|
||||
message=message,
|
||||
detail=detail,
|
||||
**common,
|
||||
)
|
||||
return
|
||||
if response_data is not None:
|
||||
log_ai_model_event(
|
||||
event_type=event_type,
|
||||
event_phase="RESPONSE",
|
||||
event_status=event_status,
|
||||
response=response_data,
|
||||
token_usage=token_usage,
|
||||
message=message,
|
||||
detail=detail,
|
||||
error=error if normalized_status == str(LogEventStatusEnum.FAILED.value).lower() else None,
|
||||
**common,
|
||||
)
|
||||
if normalized_status == str(LogEventStatusEnum.FAILED.value).lower() or error:
|
||||
log_ai_model_event(
|
||||
event_type=event_type,
|
||||
event_phase="ERROR",
|
||||
event_status=LogEventStatusEnum.FAILED.value,
|
||||
message=message,
|
||||
detail=detail,
|
||||
error=error or "AI model call failed",
|
||||
**common,
|
||||
)
|
||||
|
||||
async def _select_model_config(db: AsyncSession) -> ModelConfig | None:
|
||||
result = await db.execute(select(ModelConfig).where(ModelConfig.is_active == True, ModelConfig.deleted_at.is_(None)).order_by(ModelConfig.priority.desc()).limit(1))
|
||||
@@ -1536,6 +1568,7 @@ async def optimize_hot_opening_video_prompt(
|
||||
step_id: str | None = None,
|
||||
trace_id: str | None = None,
|
||||
) -> tuple[dict[str, Any], str, dict[str, Any]]:
|
||||
call_id = generate_id()
|
||||
duration = int(video_config["duration"])
|
||||
from app.utils.media import media_to_base64, get_llm_media_as_base64
|
||||
use_base64 = await get_llm_media_as_base64(db)
|
||||
@@ -1559,9 +1592,22 @@ async def optimize_hot_opening_video_prompt(
|
||||
# result = ensure_negative_prompt(ensure_flow_matches_time_plan(ensure_top_keys(fill_none_with_wu(result)), duration))
|
||||
# return result, build_final_video_prompt(result), {"input_tokens": 0, "output_tokens": 0, "total_tokens": 0}
|
||||
|
||||
config = await _select_model_config(db)
|
||||
config_row = await _select_model_config(db)
|
||||
config = (
|
||||
SimpleNamespace(
|
||||
id=str(config_row.id),
|
||||
name=str(config_row.name or ""),
|
||||
provider=str(config_row.provider or ""),
|
||||
api_base=str(config_row.api_base or ""),
|
||||
api_key=str(config_row.api_key or ""),
|
||||
model_name=str(config_row.model_name or ""),
|
||||
)
|
||||
if config_row is not None
|
||||
else None
|
||||
)
|
||||
# All module/project claims are committed by the caller. Release this
|
||||
# configuration read transaction before the remote model request.
|
||||
# configuration read transaction before the remote model request and use
|
||||
# only the scalar snapshot afterwards.
|
||||
await db.commit()
|
||||
if not config:
|
||||
result = normalize_video_prompt_schema_from_ai(_mock_result(video_config, target_platform), video_config, schema_config_snapshot)
|
||||
@@ -1602,6 +1648,7 @@ async def optimize_hot_opening_video_prompt(
|
||||
}
|
||||
started_event, remote_action = _video_prompt_remote_event(module, started=True)
|
||||
_log_video_prompt_ai_event(
|
||||
call_id=call_id,
|
||||
module=module,
|
||||
event_type=started_event,
|
||||
action=remote_action,
|
||||
@@ -1625,6 +1672,7 @@ async def optimize_hot_opening_video_prompt(
|
||||
except Exception as exc:
|
||||
failed_event, remote_action = _video_prompt_remote_event(module)
|
||||
_log_video_prompt_ai_event(
|
||||
call_id=call_id,
|
||||
module=module,
|
||||
event_type=failed_event,
|
||||
action=remote_action,
|
||||
@@ -1645,6 +1693,7 @@ async def optimize_hot_opening_video_prompt(
|
||||
if response.status_code >= 400:
|
||||
failed_event, remote_action = _video_prompt_remote_event(module)
|
||||
_log_video_prompt_ai_event(
|
||||
call_id=call_id,
|
||||
module=module,
|
||||
event_type=failed_event,
|
||||
action=remote_action,
|
||||
@@ -1671,6 +1720,7 @@ async def optimize_hot_opening_video_prompt(
|
||||
except Exception as exc:
|
||||
parse_event, remote_action = _video_prompt_remote_event(module, empty="content 为空" in str(exc), parse_failed="content 为空" not in str(exc))
|
||||
_log_video_prompt_ai_event(
|
||||
call_id=call_id,
|
||||
module=module,
|
||||
event_type=parse_event,
|
||||
action=remote_action,
|
||||
@@ -1721,6 +1771,7 @@ async def optimize_hot_opening_video_prompt(
|
||||
except Exception as exc:
|
||||
parse_event, remote_action = _video_prompt_remote_event(module, parse_failed=True)
|
||||
_log_video_prompt_ai_event(
|
||||
call_id=call_id,
|
||||
module=module,
|
||||
event_type=parse_event,
|
||||
action=remote_action,
|
||||
@@ -1740,6 +1791,7 @@ async def optimize_hot_opening_video_prompt(
|
||||
raise
|
||||
success_event, remote_action = _video_prompt_remote_event(module, success=True)
|
||||
_log_video_prompt_ai_event(
|
||||
call_id=call_id,
|
||||
module=module,
|
||||
event_type=success_event,
|
||||
action=remote_action,
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
from datetime import datetime
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
@@ -16,7 +16,8 @@ from app.enums.generation_provider import (
|
||||
)
|
||||
from app.enums.private_portrait import PRIVATE_PORTRAIT_ASSET_URI_PREFIX
|
||||
from app.models.image_engine import ImageEngine
|
||||
from app.services.log_config import LOG_DATE_FORMAT, LOG_DIR, encrypt_data, is_enabled
|
||||
from app.services.operation_log_service import build_exception_detail, log_ai_model_event
|
||||
from app.utils.id_gen import generate_id
|
||||
from app.types.generation.provider import (
|
||||
ImageProviderBatchResult,
|
||||
ImageProviderItem,
|
||||
@@ -59,49 +60,27 @@ class ImageProviderError(RuntimeError):
|
||||
}
|
||||
|
||||
|
||||
def _log_image_request(engine: ProviderImageEngineLike, record_id: str, request_data: dict):
|
||||
if not is_enabled():
|
||||
return
|
||||
try:
|
||||
os.makedirs(LOG_DIR, exist_ok=True)
|
||||
today = datetime.now().strftime(LOG_DATE_FORMAT)
|
||||
log_file = os.path.join(LOG_DIR, f"{today}.log")
|
||||
request_str = json.dumps(request_data, ensure_ascii=False)
|
||||
request_encrypted = encrypt_data(request_data, True)
|
||||
entry = {
|
||||
"timestamp": datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
|
||||
"type": "image_gen_request",
|
||||
"engine": engine.name,
|
||||
"model": engine.model_name,
|
||||
"record_id": record_id,
|
||||
"request": request_encrypted,
|
||||
"request_length": len(request_str),
|
||||
}
|
||||
with open(log_file, "a", encoding="utf-8") as file:
|
||||
file.write(json.dumps(entry, ensure_ascii=False) + "\n")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _log_image_response(record_id: str, response_data: dict, error: str | None = None):
|
||||
if not is_enabled():
|
||||
return
|
||||
try:
|
||||
os.makedirs(LOG_DIR, exist_ok=True)
|
||||
today = datetime.now().strftime(LOG_DATE_FORMAT)
|
||||
log_file = os.path.join(LOG_DIR, f"{today}.log")
|
||||
response_encrypted = encrypt_data(response_data, True) if response_data else ""
|
||||
entry = {
|
||||
"timestamp": datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
|
||||
"type": "image_gen_response",
|
||||
"record_id": record_id,
|
||||
"response": response_encrypted,
|
||||
"error": error,
|
||||
}
|
||||
with open(log_file, "a", encoding="utf-8") as file:
|
||||
file.write(json.dumps(entry, ensure_ascii=False) + "\n")
|
||||
except Exception:
|
||||
pass
|
||||
def _provider_log_context(engine, record, *, call_id: str, step_code: str) -> dict:
|
||||
generation_mode = str(getattr(record, "generation_mode", "") or "generation_record")
|
||||
owner_type = "chat_generation_task" if generation_mode != "generation_record" else "generation_record"
|
||||
return {
|
||||
"module": generation_mode,
|
||||
"step_code": step_code,
|
||||
"call_id": call_id,
|
||||
"source": "app.services.image_gen",
|
||||
"user_id": str(getattr(record, "user_id", "") or "") or None,
|
||||
"project_id": str(getattr(record, "project_id", "") or "") or None,
|
||||
"task_id": str(getattr(record, "id", "") or "") or None,
|
||||
"owner_type": owner_type,
|
||||
"owner_id": str(getattr(record, "id", "") or "") or None,
|
||||
"generation_attempt_no": int(getattr(record, "generation_attempt_no", 1) or 1),
|
||||
"model_config_id": str(getattr(engine, "id", "") or "") or None,
|
||||
"model_config_name": str(getattr(engine, "name", "") or "") or None,
|
||||
"model_name": str(getattr(engine, "model_name", "") or "") or None,
|
||||
"provider": str(getattr(engine, "provider", "") or "") or None,
|
||||
"api_base": str(getattr(engine, "api_base", "") or "") or None,
|
||||
}
|
||||
|
||||
|
||||
async def get_active_image_engine(db: AsyncSession) -> ImageEngine:
|
||||
@@ -321,7 +300,18 @@ def submit_image_task(
|
||||
)
|
||||
request_sdk_payload["stream"] = False
|
||||
|
||||
_log_image_request(engine, record.id, request_log_payload)
|
||||
call_id = generate_id()
|
||||
started = time.perf_counter()
|
||||
api_step = "image_sync_batch_create" if count > 1 else "image_sync_create"
|
||||
log_context = _provider_log_context(engine, record, call_id=call_id, step_code=api_step)
|
||||
log_ai_model_event(
|
||||
event_type="REQUEST",
|
||||
event_phase="REQUEST",
|
||||
event_status="started",
|
||||
remote_action=api_step,
|
||||
request=request_log_payload,
|
||||
**log_context,
|
||||
)
|
||||
|
||||
try:
|
||||
result = client.images.generate(**request_sdk_payload)
|
||||
@@ -386,7 +376,16 @@ def submit_image_task(
|
||||
"total_tokens": total_tokens,
|
||||
},
|
||||
}
|
||||
_log_image_response(record.id, response_data)
|
||||
log_ai_model_event(
|
||||
event_type="RESPONSE",
|
||||
event_phase="RESPONSE",
|
||||
event_status="success",
|
||||
remote_action=api_step,
|
||||
latency_ms=int((time.perf_counter() - started) * 1000),
|
||||
response=response_data,
|
||||
token_usage=response_data.get("usage"),
|
||||
**log_context,
|
||||
)
|
||||
return {
|
||||
"items": items,
|
||||
"model": str(response_data["model"] or ""),
|
||||
@@ -404,7 +403,18 @@ def submit_image_task(
|
||||
provider_error.error_code,
|
||||
provider_error.safe_message,
|
||||
)
|
||||
_log_image_response(record.id, provider_error.as_dict(), provider_error.safe_message)
|
||||
log_ai_model_event(
|
||||
event_type="ERROR",
|
||||
event_phase="ERROR",
|
||||
event_status="failed",
|
||||
remote_action=api_step,
|
||||
http_status=provider_error.http_status,
|
||||
remote_request_id=provider_error.provider_request_id,
|
||||
latency_ms=int((time.perf_counter() - started) * 1000),
|
||||
detail=build_exception_detail(exc, provider_error.as_dict()),
|
||||
error=provider_error.safe_message,
|
||||
**log_context,
|
||||
)
|
||||
raise provider_error from exc
|
||||
finally:
|
||||
try:
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import json
|
||||
import os
|
||||
from datetime import datetime
|
||||
import time
|
||||
from types import SimpleNamespace
|
||||
|
||||
import httpx
|
||||
from sqlalchemy import select
|
||||
@@ -10,47 +10,13 @@ from app.config import settings
|
||||
from app.models.model_config import ModelConfig
|
||||
from app.models.token_usage import TokenUsage
|
||||
from app.utils.id_gen import generate_id
|
||||
from app.services.log_config import is_enabled, LOG_DIR, LOG_DATE_FORMAT, encrypt_data
|
||||
from app.services.operation_log_service import build_exception_detail, log_ai_model_event
|
||||
|
||||
|
||||
def _sanitize_for_log(data):
|
||||
"""Replace base64 data URIs with placeholder for readable logs."""
|
||||
if isinstance(data, str):
|
||||
if data.startswith("data:") and ";base64," in data:
|
||||
return "[base64 image data]"
|
||||
return data
|
||||
if isinstance(data, dict):
|
||||
return {k: _sanitize_for_log(v) for k, v in data.items()}
|
||||
if isinstance(data, list):
|
||||
return [_sanitize_for_log(item) for item in data]
|
||||
return data
|
||||
|
||||
|
||||
def _log_ai_request_response(config, request_data: dict, response_data: dict | None, error: str | None = None):
|
||||
"""Log AI model request/response to log/AiModel/YYYY-MM-DD.log"""
|
||||
if not is_enabled():
|
||||
return
|
||||
try:
|
||||
os.makedirs(LOG_DIR, exist_ok=True)
|
||||
today = datetime.now().strftime(LOG_DATE_FORMAT)
|
||||
log_file = os.path.join(LOG_DIR, f"{today}.log")
|
||||
request_encrypted = encrypt_data(_sanitize_for_log(request_data), True)
|
||||
response_encrypted = encrypt_data(_sanitize_for_log(response_data), True) if response_data else ""
|
||||
entry = {
|
||||
"timestamp": datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
|
||||
"model_name": config.name,
|
||||
"model_id": config.model_name,
|
||||
"provider": config.provider,
|
||||
"api_base": config.api_base,
|
||||
"request": request_encrypted,
|
||||
"response": response_encrypted,
|
||||
"error": error,
|
||||
}
|
||||
with open(log_file, "a", encoding="utf-8") as f:
|
||||
f.write(json.dumps(entry, ensure_ascii=False) + "\n")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
class LLMProviderCallError(RuntimeError):
|
||||
"""Remote model call or response validation failed and may use fallback."""
|
||||
|
||||
MOCK_OPTIMIZED_PROMPTS = {
|
||||
"直播": "专业直播间场景,45度斜角机位,暖色柔光打光,主播居中构图,背景虚化处理,产品特写切换流畅,镜头推进节奏感强,画面色彩饱和度高,适合电商直播推广视频。",
|
||||
@@ -92,9 +58,17 @@ async def optimize_prompt(
|
||||
duration: int | None = None,
|
||||
image_size: str | None = None,
|
||||
image_proportion: str | None = None,
|
||||
image_px: str | None | None = None,
|
||||
image_px: str | None = None,
|
||||
references: list[dict] | None = None,
|
||||
gen_type: str = "video",
|
||||
*,
|
||||
log_module: str = "generation_ai",
|
||||
log_step: str = "prompt_optimize",
|
||||
log_project_id: str | None = None,
|
||||
log_task_id: str | None = None,
|
||||
log_owner_type: str | None = None,
|
||||
log_owner_id: str | None = None,
|
||||
generation_attempt_no: int | None = None,
|
||||
) -> tuple[str, dict]:
|
||||
"""Optimize user prompt using LLM. Returns (optimized_text, token_usage_dict)."""
|
||||
|
||||
@@ -103,9 +77,22 @@ async def optimize_prompt(
|
||||
.where(ModelConfig.is_active == True, ModelConfig.deleted_at.is_(None))
|
||||
.order_by(ModelConfig.priority.desc())
|
||||
)
|
||||
configs = list(result.scalars().all())
|
||||
# Release the read transaction before the external LLM request. Callers
|
||||
# must commit their business claim before invoking optimize_prompt.
|
||||
configs = [
|
||||
SimpleNamespace(
|
||||
id=item.id,
|
||||
name=item.name,
|
||||
provider=item.provider,
|
||||
api_base=item.api_base,
|
||||
api_key=item.api_key,
|
||||
model_name=item.model_name,
|
||||
max_tokens=item.max_tokens,
|
||||
temperature=item.temperature,
|
||||
)
|
||||
for item in result.scalars().all()
|
||||
]
|
||||
# Release the read transaction before the external LLM request. Only
|
||||
# plain scalar snapshots are used afterwards, so expire_on_commit does
|
||||
# not trigger an ORM refresh while the provider request is in flight.
|
||||
await db.commit()
|
||||
|
||||
if configs:
|
||||
@@ -122,8 +109,15 @@ async def optimize_prompt(
|
||||
image_size=image_size,
|
||||
image_proportion=image_proportion,
|
||||
image_px=image_px,
|
||||
log_module=log_module,
|
||||
log_step=log_step,
|
||||
log_project_id=log_project_id,
|
||||
log_task_id=log_task_id,
|
||||
log_owner_type=log_owner_type,
|
||||
log_owner_id=log_owner_id,
|
||||
generation_attempt_no=generation_attempt_no,
|
||||
)
|
||||
except Exception:
|
||||
except LLMProviderCallError:
|
||||
continue
|
||||
|
||||
# 所有真实模型都失败,降级到 mock
|
||||
@@ -156,7 +150,15 @@ async def _call_openai_compatible(
|
||||
gen_type: str = "video",
|
||||
image_size: str | None = None,
|
||||
image_proportion: str | None = None,
|
||||
image_px: str | None | None = None,
|
||||
image_px: str | None = None,
|
||||
*,
|
||||
log_module: str = "generation_ai",
|
||||
log_step: str = "prompt_optimize",
|
||||
log_project_id: str | None = None,
|
||||
log_task_id: str | None = None,
|
||||
log_owner_type: str | None = None,
|
||||
log_owner_id: str | None = None,
|
||||
generation_attempt_no: int | None = None,
|
||||
) -> tuple[str, dict]:
|
||||
"""Call an OpenAI-compatible API to optimize the prompt. Returns (content, token_usage)."""
|
||||
system_prompt = None
|
||||
@@ -321,14 +323,33 @@ async def _call_openai_compatible(
|
||||
"max_tokens": config.max_tokens,
|
||||
"temperature": config.temperature,
|
||||
}
|
||||
# Build log-friendly request data (image paths instead of base64)
|
||||
if log_user_message:
|
||||
log_request_data = {**request_data, "messages": [
|
||||
{"role": "system", "content": system_prompt},
|
||||
log_user_message,
|
||||
]}
|
||||
else:
|
||||
log_request_data = request_data
|
||||
call_id = generate_id()
|
||||
started = time.perf_counter()
|
||||
common_log = {
|
||||
"module": log_module,
|
||||
"step_code": log_step,
|
||||
"call_id": call_id,
|
||||
"source": "app.services.llm",
|
||||
"user_id": user_id,
|
||||
"project_id": log_project_id,
|
||||
"task_id": log_task_id,
|
||||
"owner_type": log_owner_type,
|
||||
"owner_id": log_owner_id,
|
||||
"generation_attempt_no": generation_attempt_no,
|
||||
"model_config_id": config.id,
|
||||
"model_config_name": config.name,
|
||||
"model_name": config.model_name,
|
||||
"provider": config.provider,
|
||||
"api_base": config.api_base,
|
||||
"remote_action": "chat_completions",
|
||||
}
|
||||
log_ai_model_event(
|
||||
event_type="REQUEST",
|
||||
event_phase="REQUEST",
|
||||
event_status="started",
|
||||
request=request_data,
|
||||
**common_log,
|
||||
)
|
||||
try:
|
||||
response = await client.post(
|
||||
f"{config.api_base}/chat/completions",
|
||||
@@ -338,47 +359,117 @@ async def _call_openai_compatible(
|
||||
},
|
||||
json=request_data,
|
||||
)
|
||||
latency_ms = int((time.perf_counter() - started) * 1000)
|
||||
if response.status_code >= 400:
|
||||
error_body = response.text
|
||||
_log_ai_request_response(config, log_request_data, None, error=f"HTTP {response.status_code}: {error_body}")
|
||||
raise RuntimeError(f"HTTP {response.status_code}: {error_body}")
|
||||
log_ai_model_event(
|
||||
event_type="RESPONSE",
|
||||
event_phase="RESPONSE",
|
||||
event_status="failed",
|
||||
http_status=response.status_code,
|
||||
latency_ms=latency_ms,
|
||||
response={"body": error_body},
|
||||
error=f"HTTP {response.status_code}",
|
||||
**common_log,
|
||||
)
|
||||
error = LLMProviderCallError(f"HTTP {response.status_code}: {error_body}")
|
||||
log_ai_model_event(
|
||||
event_type="ERROR",
|
||||
event_phase="ERROR",
|
||||
event_status="failed",
|
||||
http_status=response.status_code,
|
||||
latency_ms=latency_ms,
|
||||
detail=build_exception_detail(error),
|
||||
error=str(error),
|
||||
**common_log,
|
||||
)
|
||||
raise error
|
||||
data = response.json()
|
||||
except RuntimeError:
|
||||
log_ai_model_event(
|
||||
event_type="RESPONSE",
|
||||
event_phase="RESPONSE",
|
||||
event_status="success",
|
||||
http_status=response.status_code,
|
||||
latency_ms=latency_ms,
|
||||
response=data,
|
||||
token_usage=data.get("usage") if isinstance(data, dict) else None,
|
||||
**common_log,
|
||||
)
|
||||
except LLMProviderCallError:
|
||||
raise
|
||||
except Exception as e:
|
||||
_log_ai_request_response(config, log_request_data, None, error=str(e))
|
||||
raise RuntimeError(f"{type(e).__name__}: {e}")
|
||||
except Exception as exc:
|
||||
latency_ms = int((time.perf_counter() - started) * 1000)
|
||||
log_ai_model_event(
|
||||
event_type="ERROR",
|
||||
event_phase="ERROR",
|
||||
event_status="failed",
|
||||
latency_ms=latency_ms,
|
||||
detail=build_exception_detail(exc),
|
||||
error=str(exc),
|
||||
**common_log,
|
||||
)
|
||||
raise LLMProviderCallError(f"{type(exc).__name__}: {exc}") from exc
|
||||
|
||||
# Log request/response
|
||||
_log_ai_request_response(config, log_request_data, data)
|
||||
|
||||
# Record token usage
|
||||
usage = data.get("usage", {})
|
||||
input_tokens = usage.get("prompt_tokens", 0)
|
||||
output_tokens = usage.get("completion_tokens", 0)
|
||||
total_tokens = usage.get("total_tokens", input_tokens + output_tokens)
|
||||
try:
|
||||
usage = data.get("usage", {})
|
||||
input_tokens = int(usage.get("prompt_tokens", 0) or 0)
|
||||
output_tokens = int(usage.get("completion_tokens", 0) or 0)
|
||||
total_tokens = int(usage.get("total_tokens", input_tokens + output_tokens) or 0)
|
||||
content = data["choices"][0]["message"]["content"].strip()
|
||||
if not content:
|
||||
raise ValueError("模型未返回有效提示词")
|
||||
except Exception as exc:
|
||||
log_ai_model_event(
|
||||
event_type="ERROR",
|
||||
event_phase="ERROR",
|
||||
event_status="failed",
|
||||
latency_ms=int((time.perf_counter() - started) * 1000),
|
||||
detail=build_exception_detail(exc, {"stage": "response_validation"}),
|
||||
error=str(exc),
|
||||
**common_log,
|
||||
)
|
||||
raise LLMProviderCallError(f"模型响应解析失败: {exc}") from exc
|
||||
|
||||
token_usage_id = None
|
||||
if db is not None:
|
||||
token_usage_id = generate_id()
|
||||
record = TokenUsage(
|
||||
id=token_usage_id,
|
||||
model_config_id=config.id,
|
||||
user_id=user_id,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
total_tokens=total_tokens,
|
||||
)
|
||||
db.add(record)
|
||||
await db.flush()
|
||||
try:
|
||||
token_usage_id = generate_id()
|
||||
record = TokenUsage(
|
||||
id=token_usage_id,
|
||||
model_config_id=config.id,
|
||||
user_id=user_id,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
total_tokens=total_tokens,
|
||||
source_module=log_module,
|
||||
source_step_code=log_step,
|
||||
owner_type=log_owner_type,
|
||||
owner_id=log_owner_id,
|
||||
)
|
||||
db.add(record)
|
||||
await db.flush()
|
||||
except Exception as exc:
|
||||
log_ai_model_event(
|
||||
event_type="ERROR",
|
||||
event_phase="ERROR",
|
||||
event_status="failed",
|
||||
latency_ms=int((time.perf_counter() - started) * 1000),
|
||||
detail=build_exception_detail(exc, {"stage": "token_usage_persistence"}),
|
||||
error=str(exc),
|
||||
**common_log,
|
||||
)
|
||||
# A local transaction failure must not call a second provider after
|
||||
# the first provider has already returned a valid response.
|
||||
raise
|
||||
|
||||
content = data["choices"][0]["message"]["content"].strip()
|
||||
token_usage = {
|
||||
"token_usage_id": token_usage_id,
|
||||
"model_config_id": config.id,
|
||||
"model_config_name": config.name,
|
||||
"model_provider": config.provider,
|
||||
"model_name": config.model_name,
|
||||
"source_module": log_module,
|
||||
"source_step_code": log_step,
|
||||
"input_tokens": input_tokens,
|
||||
"output_tokens": output_tokens,
|
||||
"total_tokens": total_tokens,
|
||||
|
||||
@@ -61,8 +61,13 @@ def log_module_prompt_event(
|
||||
step_id=step_id,
|
||||
user_id=user_id,
|
||||
trace_id=trace_id,
|
||||
message=f"模块 AI 请求:{prompt_type}",
|
||||
detail={"prompt_type": prompt_type, "request": request or {}, "response": response or {}, "token_usage": token_usage or {}},
|
||||
message=f"模块 AI 步骤:{prompt_type}",
|
||||
detail={
|
||||
"prompt_type": prompt_type,
|
||||
"has_request": request is not None,
|
||||
"has_response": response is not None,
|
||||
"token_usage": token_usage or {},
|
||||
},
|
||||
error=error,
|
||||
event_status=LogEventStatusEnum.FAILED.value if error else LogEventStatusEnum.SUCCESS.value,
|
||||
source=LogSourceEnum.SERVICE.value,
|
||||
|
||||
@@ -21,7 +21,6 @@ MODULE_GENERATION_LOG_ROOT = os.path.join(LOG_BASE_DIR, "ModuleGeneration")
|
||||
AI_MODEL_LOG_ROOT = LOG_DIR
|
||||
SENSITIVE_KEY_PATTERNS = (
|
||||
"secret",
|
||||
"token",
|
||||
"authorization",
|
||||
"cookie",
|
||||
"credential",
|
||||
@@ -30,9 +29,28 @@ SENSITIVE_KEY_PATTERNS = (
|
||||
"access_key",
|
||||
"api_key",
|
||||
"apikey",
|
||||
"security-token",
|
||||
"x-tos-security-token",
|
||||
"security_token",
|
||||
)
|
||||
SENSITIVE_TOKEN_KEYS = {
|
||||
"token",
|
||||
"access_token",
|
||||
"refresh_token",
|
||||
"bearer_token",
|
||||
"security_token",
|
||||
"x_tos_security_token",
|
||||
}
|
||||
FILE_BASE64_KEYS = {
|
||||
"b64_json",
|
||||
"file_data",
|
||||
"file_base64",
|
||||
"content_base64",
|
||||
"image_base64",
|
||||
"video_base64",
|
||||
"audio_base64",
|
||||
}
|
||||
FILE_DATA_URI_MIME_PREFIXES = ("image/", "video/", "audio/")
|
||||
FILE_DATA_URI_MIME_TYPES = {"application/pdf", "application/octet-stream"}
|
||||
FILE_BASE64_PREVIEW_CHARS = 30
|
||||
|
||||
|
||||
def _safe_name(value: str | None, default: str = "unknown") -> str:
|
||||
@@ -49,7 +67,41 @@ def _mask_string(value: str) -> str:
|
||||
|
||||
def _is_sensitive_key(key: str) -> bool:
|
||||
lower = str(key).replace("-", "_").lower()
|
||||
return lower == "sign" or any(pattern in lower for pattern in SENSITIVE_KEY_PATTERNS)
|
||||
if lower == "sign" or lower in SENSITIVE_TOKEN_KEYS:
|
||||
return True
|
||||
return any(pattern in lower for pattern in SENSITIVE_KEY_PATTERNS)
|
||||
|
||||
|
||||
def _decoded_base64_size(value: str) -> int:
|
||||
compact = "".join(value.split())
|
||||
if not compact:
|
||||
return 0
|
||||
padding = 2 if compact.endswith("==") else (1 if compact.endswith("=") else 0)
|
||||
return max(0, (len(compact) * 3) // 4 - padding)
|
||||
|
||||
|
||||
def _file_base64_preview(value: str, key_path: tuple[str, ...]) -> str | None:
|
||||
data_uri = re.match(r"^data:([^;,]+);base64,(.*)$", value, flags=re.IGNORECASE | re.DOTALL)
|
||||
prefix = ""
|
||||
payload = value
|
||||
is_file = False
|
||||
if data_uri:
|
||||
mime = str(data_uri.group(1) or "").lower()
|
||||
is_file = mime.startswith(FILE_DATA_URI_MIME_PREFIXES) or mime in FILE_DATA_URI_MIME_TYPES
|
||||
prefix = value[: value.find(",") + 1]
|
||||
payload = data_uri.group(2)
|
||||
elif key_path and key_path[-1].replace("-", "_").lower() in FILE_BASE64_KEYS:
|
||||
# Raw base64 is treated as file content only for an explicit file field.
|
||||
is_file = len(value) >= 64 and bool(re.fullmatch(r"[A-Za-z0-9+/=\s]+", value))
|
||||
if not is_file:
|
||||
return None
|
||||
|
||||
compact = "".join(payload.split())
|
||||
preview = compact[:FILE_BASE64_PREVIEW_CHARS]
|
||||
total_bytes = _decoded_base64_size(compact)
|
||||
preview_bytes = min(total_bytes, (len(preview) * 3) // 4)
|
||||
remaining_bytes = max(0, total_bytes - preview_bytes)
|
||||
return f"{prefix}{preview}...<remaining_file_bytes:{remaining_bytes}>"
|
||||
|
||||
|
||||
def _sanitize_url(value: str) -> str:
|
||||
@@ -65,10 +117,13 @@ def _sanitize_url(value: str) -> str:
|
||||
return value
|
||||
|
||||
|
||||
def sanitize_log_value(value: Any) -> Any:
|
||||
def sanitize_log_value(value: Any, *, key_path: tuple[str, ...] = ()) -> Any:
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, str):
|
||||
file_preview = _file_base64_preview(value, key_path)
|
||||
if file_preview is not None:
|
||||
return file_preview
|
||||
text = _sanitize_url(value) if value.startswith(("http://", "https://")) else value
|
||||
if len(text) > MAX_LOG_FIELD_LENGTH:
|
||||
return text[:MAX_LOG_FIELD_LENGTH] + f"...<truncated:{len(text) - MAX_LOG_FIELD_LENGTH}>"
|
||||
@@ -77,10 +132,14 @@ def sanitize_log_value(value: Any) -> Any:
|
||||
output: dict[str, Any] = {}
|
||||
for k, v in value.items():
|
||||
key = str(k)
|
||||
output[key] = "***" if _is_sensitive_key(key) else sanitize_log_value(v)
|
||||
output[key] = (
|
||||
"***"
|
||||
if _is_sensitive_key(key)
|
||||
else sanitize_log_value(v, key_path=(*key_path, key))
|
||||
)
|
||||
return output
|
||||
if isinstance(value, list):
|
||||
return [sanitize_log_value(v) for v in value]
|
||||
if isinstance(value, (list, tuple)):
|
||||
return [sanitize_log_value(v, key_path=(*key_path, str(index))) for index, v in enumerate(value)]
|
||||
return value
|
||||
|
||||
|
||||
@@ -262,6 +321,13 @@ def log_ai_model_event(
|
||||
*,
|
||||
event_type: str,
|
||||
module: str | None = None,
|
||||
step_code: str | None = None,
|
||||
call_id: str | None = None,
|
||||
event_phase: str | None = None,
|
||||
owner_type: str | None = None,
|
||||
owner_id: str | None = None,
|
||||
generation_attempt_no: int | None = None,
|
||||
latency_ms: int | None = None,
|
||||
event_status: str = "success",
|
||||
source: str | None = None,
|
||||
trace_id: str | None = None,
|
||||
@@ -314,8 +380,19 @@ def log_ai_model_event(
|
||||
)
|
||||
entry.update(
|
||||
{
|
||||
"call_id": call_id,
|
||||
"event_phase": event_phase or event_type,
|
||||
"step_code": step_code,
|
||||
"owner_type": owner_type,
|
||||
"owner_id": owner_id,
|
||||
"generation_attempt_no": generation_attempt_no,
|
||||
"latency_ms": latency_ms,
|
||||
# Preserve the legacy fields for existing log readers while also
|
||||
# exposing unambiguous configuration/provider model names.
|
||||
"model_name": model_config_name,
|
||||
"model_id": model_name,
|
||||
"model_config_name": model_config_name,
|
||||
"provider_model_name": model_name,
|
||||
"model_config_id": model_config_id,
|
||||
"provider": provider,
|
||||
"api_base": api_base,
|
||||
|
||||
@@ -2,6 +2,8 @@ from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import time
|
||||
import uuid
|
||||
from typing import Any
|
||||
|
||||
from fastapi import HTTPException
|
||||
@@ -14,10 +16,8 @@ from app.enums.private_portrait import (
|
||||
ARK_PRIVATE_PORTRAIT_VERSION,
|
||||
ArkPrivatePortraitAction,
|
||||
PrivatePortraitEventSource,
|
||||
PrivatePortraitEventStatus,
|
||||
PrivatePortraitEventType,
|
||||
)
|
||||
from app.services.operation_log_service import log_remote_api_event
|
||||
from app.services.operation_log_service import log_ai_model_event
|
||||
from app.services.private_portrait.rate_limiter import acquire_private_portrait_action_token
|
||||
|
||||
DOMAIN = "private_portrait"
|
||||
@@ -117,39 +117,69 @@ class ArkPrivateAssetClient:
|
||||
async def _call(self, action: ArkPrivatePortraitAction, payload: dict[str, Any]) -> dict[str, Any]:
|
||||
action_value = action.value
|
||||
await acquire_private_portrait_action_token(action=action_value, wait_timeout_seconds=2.0, for_celery=self.for_celery)
|
||||
log_remote_api_event(
|
||||
domain=DOMAIN,
|
||||
call_id = uuid.uuid4().hex
|
||||
source = PrivatePortraitEventSource.CELERY.value if self.for_celery else PrivatePortraitEventSource.SERVICE.value
|
||||
started = time.perf_counter()
|
||||
log_ai_model_event(
|
||||
event_type="REQUEST",
|
||||
event_phase="REQUEST",
|
||||
event_status="started",
|
||||
module=DOMAIN,
|
||||
step_code=action_value,
|
||||
call_id=call_id,
|
||||
source=source,
|
||||
remote_action=action_value,
|
||||
event_type=PrivatePortraitEventType.ARK_API_CALL_START.value,
|
||||
event_status=PrivatePortraitEventStatus.PENDING.value,
|
||||
source=PrivatePortraitEventSource.CELERY.value if self.for_celery else PrivatePortraitEventSource.SERVICE.value,
|
||||
provider="volcengine_ark",
|
||||
request=payload,
|
||||
)
|
||||
try:
|
||||
result = await asyncio.to_thread(self._call_sync, action, payload)
|
||||
log_remote_api_event(
|
||||
domain=DOMAIN,
|
||||
log_ai_model_event(
|
||||
event_type="RESPONSE",
|
||||
event_phase="RESPONSE",
|
||||
event_status="success",
|
||||
module=DOMAIN,
|
||||
step_code=action_value,
|
||||
call_id=call_id,
|
||||
source=source,
|
||||
remote_action=action_value,
|
||||
event_type=PrivatePortraitEventType.ARK_API_CALL_SUCCESS.value,
|
||||
event_status=PrivatePortraitEventStatus.SUCCESS.value,
|
||||
source=PrivatePortraitEventSource.CELERY.value if self.for_celery else PrivatePortraitEventSource.SERVICE.value,
|
||||
request=payload,
|
||||
response=result,
|
||||
remote_request_id=result.get("RequestId") or result.get("request_id"),
|
||||
provider="volcengine_ark",
|
||||
latency_ms=int((time.perf_counter() - started) * 1000),
|
||||
response=result,
|
||||
)
|
||||
return result
|
||||
except ArkPrivateAssetRemoteError as exc:
|
||||
log_remote_api_event(
|
||||
domain=DOMAIN,
|
||||
log_ai_model_event(
|
||||
event_type="RESPONSE",
|
||||
event_phase="RESPONSE",
|
||||
event_status="failed",
|
||||
module=DOMAIN,
|
||||
step_code=action_value,
|
||||
call_id=call_id,
|
||||
source=source,
|
||||
remote_action=action_value,
|
||||
event_type=PrivatePortraitEventType.ARK_API_CALL_FAILED.value,
|
||||
event_status=PrivatePortraitEventStatus.FAILED.value,
|
||||
source=PrivatePortraitEventSource.CELERY.value if self.for_celery else PrivatePortraitEventSource.SERVICE.value,
|
||||
request=payload,
|
||||
response=exc.raw,
|
||||
remote_request_id=exc.request_id,
|
||||
remote_code=exc.code,
|
||||
remote_message=exc.message,
|
||||
provider="volcengine_ark",
|
||||
latency_ms=int((time.perf_counter() - started) * 1000),
|
||||
response=exc.raw,
|
||||
detail={"remote_code": exc.code, "remote_message": exc.message},
|
||||
error=exc.message,
|
||||
)
|
||||
log_ai_model_event(
|
||||
event_type="ERROR",
|
||||
event_phase="ERROR",
|
||||
event_status="failed",
|
||||
module=DOMAIN,
|
||||
step_code=action_value,
|
||||
call_id=call_id,
|
||||
source=source,
|
||||
remote_action=action_value,
|
||||
remote_request_id=exc.request_id,
|
||||
provider="volcengine_ark",
|
||||
latency_ms=int((time.perf_counter() - started) * 1000),
|
||||
detail={"remote_code": exc.code, "remote_message": exc.message},
|
||||
error=str(exc),
|
||||
)
|
||||
if self.for_celery:
|
||||
raise
|
||||
@@ -157,14 +187,18 @@ class ArkPrivateAssetClient:
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as exc:
|
||||
log_remote_api_event(
|
||||
domain=DOMAIN,
|
||||
log_ai_model_event(
|
||||
event_type="ERROR",
|
||||
event_phase="ERROR",
|
||||
event_status="failed",
|
||||
module=DOMAIN,
|
||||
step_code=action_value,
|
||||
call_id=call_id,
|
||||
source=source,
|
||||
remote_action=action_value,
|
||||
event_type=PrivatePortraitEventType.ARK_API_CALL_FAILED.value,
|
||||
event_status=PrivatePortraitEventStatus.FAILED.value,
|
||||
source=PrivatePortraitEventSource.CELERY.value if self.for_celery else PrivatePortraitEventSource.SERVICE.value,
|
||||
request=payload,
|
||||
remote_message=str(exc),
|
||||
provider="volcengine_ark",
|
||||
latency_ms=int((time.perf_counter() - started) * 1000),
|
||||
error=str(exc),
|
||||
)
|
||||
if self.for_celery:
|
||||
raise
|
||||
|
||||
@@ -921,6 +921,13 @@ async def run_image_prompt_optimize(
|
||||
user_id=user_id_value,
|
||||
references=references,
|
||||
gen_type="image",
|
||||
log_module=module_value,
|
||||
log_step="shot_replicate_image_prompt_optimize",
|
||||
log_project_id=project_id_value,
|
||||
log_task_id=step_id_value,
|
||||
log_owner_type="module_generation_step",
|
||||
log_owner_id=step_id_value,
|
||||
generation_attempt_no=expected_step_version,
|
||||
)
|
||||
if execution_guard is not None:
|
||||
await execution_guard()
|
||||
|
||||
@@ -459,6 +459,7 @@ def _log_shot_ai_model_event(
|
||||
event_status: str,
|
||||
config: ModelConfig,
|
||||
trace_id: str,
|
||||
call_id: str,
|
||||
user_id: str | None,
|
||||
task_set_id: str | None,
|
||||
segment_id: str | None,
|
||||
@@ -487,30 +488,61 @@ def _log_shot_ai_model_event(
|
||||
"remote_message": remote_message,
|
||||
"remote_param": remote_param,
|
||||
})
|
||||
log_ai_model_event(
|
||||
event_type=event_type,
|
||||
event_status=event_status,
|
||||
source=LogSourceEnum.REMOTE_API.value,
|
||||
module=ModuleCodeEnum.SHOT_REPLICATE.value,
|
||||
trace_id=trace_id,
|
||||
user_id=user_id,
|
||||
project_id=task_set_id,
|
||||
step_id=segment_id,
|
||||
remote_action=action,
|
||||
remote_request_id=remote_request_id,
|
||||
model_config_id=str(config.id),
|
||||
model_config_name=config.name,
|
||||
model_name=config.model_name,
|
||||
provider=config.provider,
|
||||
api_base=config.api_base,
|
||||
http_status=http_status,
|
||||
request=request_data,
|
||||
response=response_data,
|
||||
token_usage=token_usage,
|
||||
message=message,
|
||||
detail=detail,
|
||||
error=error,
|
||||
)
|
||||
common = {
|
||||
"source": LogSourceEnum.REMOTE_API.value,
|
||||
"module": ModuleCodeEnum.SHOT_REPLICATE.value,
|
||||
"step_code": "source_video_analysis" if mode == "full_breakdown" else "segment_video_analysis",
|
||||
"call_id": call_id,
|
||||
"trace_id": trace_id,
|
||||
"user_id": user_id,
|
||||
"project_id": task_set_id,
|
||||
"task_id": segment_id or task_set_id,
|
||||
"step_id": segment_id,
|
||||
"owner_type": "shot_replicate_segment" if segment_id else "shot_replicate_task_set",
|
||||
"owner_id": segment_id or task_set_id,
|
||||
"remote_action": action,
|
||||
"remote_request_id": remote_request_id,
|
||||
"model_config_id": str(config.id),
|
||||
"model_config_name": config.name,
|
||||
"model_name": config.model_name,
|
||||
"provider": config.provider,
|
||||
"api_base": config.api_base,
|
||||
"http_status": http_status,
|
||||
}
|
||||
normalized_status = str(event_status or "").lower()
|
||||
if normalized_status == str(LogEventStatusEnum.STARTED.value).lower():
|
||||
log_ai_model_event(
|
||||
event_type=event_type,
|
||||
event_phase="REQUEST",
|
||||
event_status=event_status,
|
||||
request=request_data,
|
||||
message=message,
|
||||
detail=detail,
|
||||
**common,
|
||||
)
|
||||
return
|
||||
if response_data is not None:
|
||||
log_ai_model_event(
|
||||
event_type=event_type,
|
||||
event_phase="RESPONSE",
|
||||
event_status=event_status,
|
||||
response=response_data,
|
||||
token_usage=token_usage,
|
||||
message=message,
|
||||
detail=detail,
|
||||
error=error if normalized_status == str(LogEventStatusEnum.FAILED.value).lower() else None,
|
||||
**common,
|
||||
)
|
||||
if normalized_status == str(LogEventStatusEnum.FAILED.value).lower() or error:
|
||||
log_ai_model_event(
|
||||
event_type=event_type,
|
||||
event_phase="ERROR",
|
||||
event_status=LogEventStatusEnum.FAILED.value,
|
||||
message=message,
|
||||
detail=detail,
|
||||
error=error or remote_message or "AI model call failed",
|
||||
**common,
|
||||
)
|
||||
|
||||
async def analyze_video_for_shot_split(
|
||||
db: AsyncSession,
|
||||
@@ -529,6 +561,7 @@ async def analyze_video_for_shot_split(
|
||||
也不再 fallback 到 SEEDANCE_*,避免拆镜分析走错通道。
|
||||
"""
|
||||
trace_id = trace_id or generate_id()
|
||||
call_id = generate_id()
|
||||
config_row = await _select_model_config(db)
|
||||
if not config_row:
|
||||
raise RuntimeError("拆镜分析模型未配置:请先在 model_configs 表启用可用模型")
|
||||
@@ -585,6 +618,7 @@ async def analyze_video_for_shot_split(
|
||||
await db.rollback()
|
||||
|
||||
_log_shot_ai_model_event(
|
||||
call_id=call_id,
|
||||
event_type=(
|
||||
ShotReplicateLogEventEnum.ANALYSIS_REMOTE_API_STARTED.value
|
||||
if mode == "full_breakdown"
|
||||
@@ -610,6 +644,7 @@ async def analyze_video_for_shot_split(
|
||||
)
|
||||
except Exception as exc:
|
||||
_log_shot_ai_model_event(
|
||||
call_id=call_id,
|
||||
event_type=(
|
||||
ShotReplicateLogEventEnum.ANALYSIS_REMOTE_API_FAILED.value
|
||||
if mode == "full_breakdown"
|
||||
@@ -633,6 +668,7 @@ async def analyze_video_for_shot_split(
|
||||
remote_request_id, remote_code, remote_message, remote_param = _extract_remote_error(response_data)
|
||||
if response.status_code >= 400:
|
||||
_log_shot_ai_model_event(
|
||||
call_id=call_id,
|
||||
event_type=(
|
||||
ShotReplicateLogEventEnum.ANALYSIS_REMOTE_API_FAILED.value
|
||||
if mode == "full_breakdown"
|
||||
@@ -661,6 +697,7 @@ async def analyze_video_for_shot_split(
|
||||
raw = response.json()
|
||||
except Exception as exc:
|
||||
_log_shot_ai_model_event(
|
||||
call_id=call_id,
|
||||
event_type=ShotReplicateLogEventEnum.ANALYSIS_RESPONSE_PARSE_FAILED.value,
|
||||
event_status=LogEventStatusEnum.FAILED.value,
|
||||
config=config,
|
||||
@@ -683,6 +720,7 @@ async def analyze_video_for_shot_split(
|
||||
except Exception as exc:
|
||||
event_type = ShotReplicateLogEventEnum.ANALYSIS_RESPONSE_EMPTY.value if "content 为空" in str(exc) else ShotReplicateLogEventEnum.ANALYSIS_RESPONSE_PARSE_FAILED.value
|
||||
_log_shot_ai_model_event(
|
||||
call_id=call_id,
|
||||
event_type=event_type,
|
||||
event_status=LogEventStatusEnum.FAILED.value,
|
||||
config=config,
|
||||
@@ -718,7 +756,6 @@ async def analyze_video_for_shot_split(
|
||||
"split_max_seconds": _split_max_seconds(),
|
||||
"analysis_mode": mode,
|
||||
"trace_id": trace_id,
|
||||
"log_request": log_request_data,
|
||||
}
|
||||
if not token_usage["total_tokens"]:
|
||||
token_usage["total_tokens"] = token_usage["input_tokens"] + token_usage["output_tokens"]
|
||||
@@ -731,6 +768,7 @@ async def analyze_video_for_shot_split(
|
||||
})
|
||||
|
||||
_log_shot_ai_model_event(
|
||||
call_id=call_id,
|
||||
event_type=(
|
||||
ShotReplicateLogEventEnum.ANALYSIS_REMOTE_API_SUCCESS.value
|
||||
if mode == "full_breakdown"
|
||||
|
||||
@@ -1,9 +1,7 @@
|
||||
import base64
|
||||
import json
|
||||
import logging
|
||||
import mimetypes
|
||||
import os
|
||||
from datetime import datetime, timezone
|
||||
import time
|
||||
|
||||
import httpx
|
||||
from sqlalchemy import select
|
||||
@@ -13,7 +11,8 @@ from volcenginesdkarkruntime import AsyncArk
|
||||
from app.config import settings
|
||||
from app.enums.private_portrait import PRIVATE_PORTRAIT_ASSET_URI_PREFIX
|
||||
from app.models.video_engine import VideoEngine
|
||||
from app.services.log_config import is_enabled, LOG_DIR, LOG_DATE_FORMAT, encrypt_data
|
||||
from app.services.operation_log_service import build_exception_detail, log_ai_model_event
|
||||
from app.utils.id_gen import generate_id
|
||||
from app.types.generation.provider import (
|
||||
ProviderGenerationRecordLike,
|
||||
ProviderVideoEngineLike,
|
||||
@@ -22,55 +21,27 @@ from app.types.generation.provider import (
|
||||
logger = logging.getLogger("videogen")
|
||||
|
||||
|
||||
def _log_video_request(engine: ProviderVideoEngineLike, record_id: str, request_data: dict):
|
||||
"""Log video generation request to log/AiModel/YYYY-MM-DD.log"""
|
||||
if not is_enabled():
|
||||
return
|
||||
try:
|
||||
os.makedirs(LOG_DIR, exist_ok=True)
|
||||
today = datetime.now().strftime(LOG_DATE_FORMAT)
|
||||
log_file = os.path.join(LOG_DIR, f"{today}.log")
|
||||
request_str = json.dumps(request_data, ensure_ascii=False)
|
||||
request_encrypted = encrypt_data(request_data, True)
|
||||
entry = {
|
||||
"timestamp": datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
|
||||
"type": "video_gen_request",
|
||||
"engine": engine.name,
|
||||
"model": engine.model_name,
|
||||
"record_id": record_id,
|
||||
"request": request_encrypted,
|
||||
"request_length": len(request_str),
|
||||
}
|
||||
with open(log_file, "a", encoding="utf-8") as f:
|
||||
f.write(json.dumps(entry, ensure_ascii=False) + "\n")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _log_video_response(record_id: str, response_data: dict, error: str | None = None):
|
||||
"""Log video generation response to log/AiModel/YYYY-MM-DD.log"""
|
||||
if not is_enabled():
|
||||
return
|
||||
try:
|
||||
os.makedirs(LOG_DIR, exist_ok=True)
|
||||
today = datetime.now().strftime(LOG_DATE_FORMAT)
|
||||
log_file = os.path.join(LOG_DIR, f"{today}.log")
|
||||
response_encrypted = encrypt_data(response_data, True) if response_data else ""
|
||||
|
||||
entry = {
|
||||
"timestamp": datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
|
||||
"type": "video_gen_response",
|
||||
"record_id": record_id,
|
||||
"response": response_encrypted,
|
||||
"error": error,
|
||||
}
|
||||
with open(log_file, "a", encoding="utf-8") as f:
|
||||
f.write(json.dumps(entry, ensure_ascii=False) + "\n")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
|
||||
def _provider_log_context(engine, record, *, call_id: str, step_code: str) -> dict:
|
||||
generation_mode = str(getattr(record, "generation_mode", "") or "generation_record")
|
||||
owner_type = "chat_generation_task" if generation_mode != "generation_record" else "generation_record"
|
||||
return {
|
||||
"module": generation_mode,
|
||||
"step_code": step_code,
|
||||
"call_id": call_id,
|
||||
"source": "app.services.video_gen",
|
||||
"user_id": str(getattr(record, "user_id", "") or "") or None,
|
||||
"project_id": str(getattr(record, "project_id", "") or "") or None,
|
||||
"task_id": str(getattr(record, "id", "") or "") or None,
|
||||
"owner_type": owner_type,
|
||||
"owner_id": str(getattr(record, "id", "") or "") or None,
|
||||
"generation_attempt_no": int(getattr(record, "generation_attempt_no", 1) or 1),
|
||||
"model_config_id": str(getattr(engine, "id", "") or "") or None,
|
||||
"model_config_name": str(getattr(engine, "name", "") or "") or None,
|
||||
"model_name": str(getattr(engine, "model_name", "") or "") or None,
|
||||
"provider": str(getattr(engine, "provider", "") or "") or None,
|
||||
"api_base": str(getattr(engine, "api_base", "") or "") or None,
|
||||
}
|
||||
|
||||
|
||||
async def get_active_engine(db: AsyncSession) -> VideoEngine:
|
||||
@@ -153,19 +124,47 @@ async def submit_video_task(
|
||||
"watermark": False,
|
||||
}
|
||||
|
||||
# Log request to AiModel log. include_media_references 只用于排查日志,不传给供应商 API。
|
||||
_log_video_request(
|
||||
call_id = generate_id()
|
||||
started = time.perf_counter()
|
||||
log_context = _provider_log_context(
|
||||
engine,
|
||||
record.id,
|
||||
{**request_payload, "include_media_references": include_media_references},
|
||||
record,
|
||||
call_id=call_id,
|
||||
step_code="video_create",
|
||||
)
|
||||
log_ai_model_event(
|
||||
event_type="REQUEST",
|
||||
event_phase="REQUEST",
|
||||
event_status="started",
|
||||
remote_action="video_create",
|
||||
request={**request_payload, "include_media_references": include_media_references},
|
||||
**log_context,
|
||||
)
|
||||
|
||||
try:
|
||||
result = await client.content_generation.tasks.create(**request_payload)
|
||||
task_id = result.id
|
||||
_log_video_response(record.id, {"task_id": task_id})
|
||||
except Exception as e:
|
||||
_log_video_response(record.id, {}, str(e))
|
||||
log_ai_model_event(
|
||||
event_type="RESPONSE",
|
||||
event_phase="RESPONSE",
|
||||
event_status="success",
|
||||
remote_action="video_create",
|
||||
remote_request_id=task_id,
|
||||
latency_ms=int((time.perf_counter() - started) * 1000),
|
||||
response={"task_id": task_id},
|
||||
**log_context,
|
||||
)
|
||||
except Exception as exc:
|
||||
log_ai_model_event(
|
||||
event_type="ERROR",
|
||||
event_phase="ERROR",
|
||||
event_status="failed",
|
||||
remote_action="video_create",
|
||||
latency_ms=int((time.perf_counter() - started) * 1000),
|
||||
detail=build_exception_detail(exc),
|
||||
error=str(exc),
|
||||
**log_context,
|
||||
)
|
||||
raise
|
||||
finally:
|
||||
await client.close()
|
||||
|
||||
@@ -6,6 +6,7 @@ import os
|
||||
import uuid
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from collections.abc import Awaitable, Callable
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
from urllib.parse import urlsplit, urlunsplit
|
||||
|
||||
@@ -305,6 +306,28 @@ async def enqueue_upscale_task(db: AsyncSession, *, upscale: VideoUpscaleTask, r
|
||||
upscale.status = VideoUpscaleTaskStatus.PENDING.value
|
||||
upscale.stage = VideoUpscaleStage.QUEUED.value
|
||||
upscale.next_retry_at = None
|
||||
log_snapshot = SimpleNamespace(
|
||||
id=upscale_id,
|
||||
chat_generation_task_id=(
|
||||
str(upscale.chat_generation_task_id)
|
||||
if upscale.chat_generation_task_id
|
||||
else None
|
||||
),
|
||||
generation_record_id=(
|
||||
str(upscale.generation_record_id)
|
||||
if upscale.generation_record_id
|
||||
else None
|
||||
),
|
||||
processor_key=processor_key,
|
||||
status=VideoUpscaleTaskStatus.PENDING.value,
|
||||
stage=VideoUpscaleStage.QUEUED.value,
|
||||
attempt_count=int(upscale.attempt_count or 0),
|
||||
failure_count=int(upscale.failure_count or 0),
|
||||
provider_task_id=(str(upscale.provider_task_id) if upscale.provider_task_id else None),
|
||||
input_source_type=upscale.input_source_type,
|
||||
target_width=upscale.target_width,
|
||||
target_height=upscale.target_height,
|
||||
)
|
||||
await db.commit()
|
||||
|
||||
try:
|
||||
@@ -324,7 +347,7 @@ async def enqueue_upscale_task(db: AsyncSession, *, upscale: VideoUpscaleTask, r
|
||||
raise RuntimeError(f"未注册的超分处理器: {processor_key}")
|
||||
log_video_upscale_event(
|
||||
event_type="upscale_task_enqueued",
|
||||
upscale_task=upscale,
|
||||
upscale_task=log_snapshot,
|
||||
detail={"reason": reason, "celery_task_id": celery_id, "processor_key": processor_key},
|
||||
)
|
||||
except Exception as exc:
|
||||
@@ -333,7 +356,7 @@ async def enqueue_upscale_task(db: AsyncSession, *, upscale: VideoUpscaleTask, r
|
||||
log_video_upscale_event(
|
||||
event_type="upscale_task_enqueue_failed",
|
||||
event_status="failed",
|
||||
upscale_task=upscale,
|
||||
upscale_task=log_snapshot,
|
||||
message=str(exc),
|
||||
detail={"reason": reason, "celery_task_id": celery_id, "processor_key": processor_key},
|
||||
error=str(exc),
|
||||
|
||||
@@ -1,12 +1,15 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
import time
|
||||
import uuid
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
from app.config import settings
|
||||
from app.enums.video_upscale import VideoUpscaleProcessorKey
|
||||
from app.services.operation_log_service import log_ai_model_event
|
||||
|
||||
|
||||
class VolcMediaKitError(RuntimeError):
|
||||
@@ -171,7 +174,24 @@ async def submit_video_enhance(
|
||||
processor=processor,
|
||||
client_token=client_token,
|
||||
)
|
||||
call_id = uuid.uuid4().hex
|
||||
started = time.perf_counter()
|
||||
log_ai_model_event(
|
||||
event_type="REQUEST",
|
||||
event_phase="REQUEST",
|
||||
event_status="started",
|
||||
module="video_upscale",
|
||||
step_code="provider_submit",
|
||||
call_id=call_id,
|
||||
source="app.services.video_upscale.volc_service",
|
||||
remote_action=endpoint,
|
||||
provider="volcengine_mediakit",
|
||||
api_base=_base_url(),
|
||||
request=payload,
|
||||
)
|
||||
timeout = max(3, int(processor.get("request_timeout_seconds") or settings.VIDEO_UPSCALE_REMOTE_REQUEST_TIMEOUT_SECONDS))
|
||||
response: httpx.Response | None = None
|
||||
data: dict[str, Any] = {}
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=timeout, follow_redirects=True) as client:
|
||||
response = await client.post(f"{_base_url()}{endpoint}", headers=_headers(), json=payload)
|
||||
@@ -179,6 +199,23 @@ async def submit_video_enhance(
|
||||
data = response.json()
|
||||
except Exception:
|
||||
data = {"success": False, "error": {"message": response.text[:2000]}}
|
||||
log_ai_model_event(
|
||||
event_type="RESPONSE",
|
||||
event_phase="RESPONSE",
|
||||
event_status="success" if response.status_code < 400 else "failed",
|
||||
module="video_upscale",
|
||||
step_code="provider_submit",
|
||||
call_id=call_id,
|
||||
source="app.services.video_upscale.volc_service",
|
||||
remote_action=endpoint,
|
||||
remote_request_id=str(data.get("request_id") or "") or None,
|
||||
provider="volcengine_mediakit",
|
||||
api_base=_base_url(),
|
||||
http_status=response.status_code,
|
||||
latency_ms=int((time.perf_counter() - started) * 1000),
|
||||
response=data,
|
||||
error=None if response.status_code < 400 else str((data.get("error") or {}).get("message") or "remote error"),
|
||||
)
|
||||
if response.status_code >= 400:
|
||||
raise _error_from_payload(
|
||||
data,
|
||||
@@ -186,18 +223,51 @@ async def submit_video_enhance(
|
||||
http_status=response.status_code,
|
||||
endpoint=endpoint,
|
||||
)
|
||||
except VolcMediaKitError:
|
||||
if not bool(data.get("success")) or not data.get("task_id"):
|
||||
raise _error_from_payload(data, "火山超分提交失败", endpoint=endpoint)
|
||||
except VolcMediaKitError as exc:
|
||||
log_ai_model_event(
|
||||
event_type="ERROR",
|
||||
event_phase="ERROR",
|
||||
event_status="failed",
|
||||
module="video_upscale",
|
||||
step_code="provider_submit",
|
||||
call_id=call_id,
|
||||
source="app.services.video_upscale.volc_service",
|
||||
remote_action=endpoint,
|
||||
remote_request_id=exc.request_id,
|
||||
provider="volcengine_mediakit",
|
||||
api_base=_base_url(),
|
||||
http_status=exc.http_status or (response.status_code if response is not None else None),
|
||||
latency_ms=int((time.perf_counter() - started) * 1000),
|
||||
detail=exc.log_detail(),
|
||||
error=str(exc),
|
||||
)
|
||||
raise
|
||||
except (httpx.TimeoutException, httpx.NetworkError) as exc:
|
||||
raise VolcMediaKitError(
|
||||
wrapped = VolcMediaKitError(
|
||||
f"火山超分提交网络异常: {exc}",
|
||||
code="NetworkError",
|
||||
retryable=True,
|
||||
endpoint=endpoint,
|
||||
) from exc
|
||||
)
|
||||
log_ai_model_event(
|
||||
event_type="ERROR",
|
||||
event_phase="ERROR",
|
||||
event_status="failed",
|
||||
module="video_upscale",
|
||||
step_code="provider_submit",
|
||||
call_id=call_id,
|
||||
source="app.services.video_upscale.volc_service",
|
||||
remote_action=endpoint,
|
||||
provider="volcengine_mediakit",
|
||||
api_base=_base_url(),
|
||||
latency_ms=int((time.perf_counter() - started) * 1000),
|
||||
detail=wrapped.log_detail(),
|
||||
error=str(wrapped),
|
||||
)
|
||||
raise wrapped from exc
|
||||
|
||||
if not bool(data.get("success")) or not data.get("task_id"):
|
||||
raise _error_from_payload(data, "火山超分提交失败", endpoint=endpoint)
|
||||
return VolcSubmitResult(
|
||||
task_id=str(data["task_id"]),
|
||||
request_id=str(data.get("request_id")) if data.get("request_id") else None,
|
||||
@@ -209,7 +279,26 @@ async def submit_video_enhance(
|
||||
|
||||
async def query_task(task_id: str, *, request_timeout_seconds: int | None = None) -> VolcQueryResult:
|
||||
endpoint = f"/api/v1/tasks/{task_id}"
|
||||
call_id = uuid.uuid4().hex
|
||||
started = time.perf_counter()
|
||||
request_payload = {"task_id": task_id}
|
||||
log_ai_model_event(
|
||||
event_type="REQUEST",
|
||||
event_phase="REQUEST",
|
||||
event_status="started",
|
||||
module="video_upscale",
|
||||
step_code="provider_poll",
|
||||
call_id=call_id,
|
||||
source="app.services.video_upscale.volc_service",
|
||||
remote_action=endpoint,
|
||||
remote_request_id=task_id,
|
||||
provider="volcengine_mediakit",
|
||||
api_base=_base_url(),
|
||||
request=request_payload,
|
||||
)
|
||||
timeout = max(3, int(request_timeout_seconds or settings.VIDEO_UPSCALE_REMOTE_REQUEST_TIMEOUT_SECONDS))
|
||||
response: httpx.Response | None = None
|
||||
data: dict[str, Any] = {}
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=timeout, follow_redirects=True) as client:
|
||||
response = await client.get(f"{_base_url()}{endpoint}", headers=_headers())
|
||||
@@ -217,6 +306,23 @@ async def query_task(task_id: str, *, request_timeout_seconds: int | None = None
|
||||
data = response.json()
|
||||
except Exception:
|
||||
data = {"success": False, "error": {"message": response.text[:2000]}}
|
||||
log_ai_model_event(
|
||||
event_type="RESPONSE",
|
||||
event_phase="RESPONSE",
|
||||
event_status="success" if response.status_code < 400 else "failed",
|
||||
module="video_upscale",
|
||||
step_code="provider_poll",
|
||||
call_id=call_id,
|
||||
source="app.services.video_upscale.volc_service",
|
||||
remote_action=endpoint,
|
||||
remote_request_id=str(data.get("request_id") or task_id),
|
||||
provider="volcengine_mediakit",
|
||||
api_base=_base_url(),
|
||||
http_status=response.status_code,
|
||||
latency_ms=int((time.perf_counter() - started) * 1000),
|
||||
response=data,
|
||||
error=None if response.status_code < 400 else str((data.get("error") or {}).get("message") or "remote error"),
|
||||
)
|
||||
if response.status_code >= 400:
|
||||
raise _error_from_payload(
|
||||
data,
|
||||
@@ -224,28 +330,62 @@ async def query_task(task_id: str, *, request_timeout_seconds: int | None = None
|
||||
http_status=response.status_code,
|
||||
endpoint=endpoint,
|
||||
)
|
||||
except VolcMediaKitError:
|
||||
if not bool(data.get("success")):
|
||||
raise _error_from_payload(data, "火山超分任务查询失败", endpoint=endpoint)
|
||||
status = str(data.get("status") or "").strip().lower()
|
||||
if status not in {"running", "completed", "failed"}:
|
||||
raise VolcMediaKitError(
|
||||
f"火山超分返回未知任务状态: {status}",
|
||||
code="UnknownStatus",
|
||||
retryable=True,
|
||||
request_id=str(data.get("request_id") or "") or None,
|
||||
endpoint=endpoint,
|
||||
response_payload=data,
|
||||
)
|
||||
except VolcMediaKitError as exc:
|
||||
log_ai_model_event(
|
||||
event_type="ERROR",
|
||||
event_phase="ERROR",
|
||||
event_status="failed",
|
||||
module="video_upscale",
|
||||
step_code="provider_poll",
|
||||
call_id=call_id,
|
||||
source="app.services.video_upscale.volc_service",
|
||||
remote_action=endpoint,
|
||||
remote_request_id=exc.request_id or task_id,
|
||||
provider="volcengine_mediakit",
|
||||
api_base=_base_url(),
|
||||
http_status=exc.http_status or (response.status_code if response is not None else None),
|
||||
latency_ms=int((time.perf_counter() - started) * 1000),
|
||||
detail=exc.log_detail(),
|
||||
error=str(exc),
|
||||
)
|
||||
raise
|
||||
except (httpx.TimeoutException, httpx.NetworkError) as exc:
|
||||
raise VolcMediaKitError(
|
||||
wrapped = VolcMediaKitError(
|
||||
f"火山超分查询网络异常: {exc}",
|
||||
code="NetworkError",
|
||||
retryable=True,
|
||||
endpoint=endpoint,
|
||||
) from exc
|
||||
|
||||
if not bool(data.get("success")):
|
||||
raise _error_from_payload(data, "火山超分任务查询失败", endpoint=endpoint)
|
||||
status = str(data.get("status") or "").strip().lower()
|
||||
if status not in {"running", "completed", "failed"}:
|
||||
raise VolcMediaKitError(
|
||||
f"火山超分返回未知任务状态: {status}",
|
||||
code="UnknownStatus",
|
||||
retryable=True,
|
||||
request_id=str(data.get("request_id") or "") or None,
|
||||
endpoint=endpoint,
|
||||
response_payload=data,
|
||||
)
|
||||
log_ai_model_event(
|
||||
event_type="ERROR",
|
||||
event_phase="ERROR",
|
||||
event_status="failed",
|
||||
module="video_upscale",
|
||||
step_code="provider_poll",
|
||||
call_id=call_id,
|
||||
source="app.services.video_upscale.volc_service",
|
||||
remote_action=endpoint,
|
||||
remote_request_id=task_id,
|
||||
provider="volcengine_mediakit",
|
||||
api_base=_base_url(),
|
||||
latency_ms=int((time.perf_counter() - started) * 1000),
|
||||
detail=wrapped.log_detail(),
|
||||
error=str(wrapped),
|
||||
)
|
||||
raise wrapped from exc
|
||||
|
||||
expires_raw = data.get("expires_at")
|
||||
try:
|
||||
expires_at = int(expires_raw) if expires_raw is not None else None
|
||||
|
||||
@@ -213,6 +213,24 @@ async def _remove_active(owner: GenerationOwner) -> None:
|
||||
await remove_download_active(_registry_id(owner))
|
||||
|
||||
|
||||
async def _reload_owner_after_commit(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
owner_type: str,
|
||||
owner_id: str,
|
||||
attempt_no: int,
|
||||
) -> GenerationOwner | None:
|
||||
owner = await load_generation_owner(
|
||||
db,
|
||||
owner_type=owner_type,
|
||||
owner_id=owner_id,
|
||||
for_update=False,
|
||||
)
|
||||
if owner is None or not is_attempt_current(owner, attempt_no):
|
||||
return None
|
||||
return owner
|
||||
|
||||
|
||||
async def _apply(
|
||||
owner: GenerationOwner,
|
||||
*,
|
||||
@@ -297,7 +315,19 @@ async def enqueue_download_task(
|
||||
if not owner.download_storage_date_dir:
|
||||
created_at = ensure_aware_utc(owner.created_at) or _now()
|
||||
owner.download_storage_date_dir = created_at.strftime("%Y/%m/%d")
|
||||
owner_type_snapshot = owner_type_of(owner)
|
||||
owner_id_snapshot = str(owner.id)
|
||||
attempt_snapshot = int(owner.generation_attempt_no or 1)
|
||||
celery_task_id_snapshot = str(owner.download_celery_task_id or "") or None
|
||||
await db.commit()
|
||||
owner = await load_generation_owner(
|
||||
db,
|
||||
owner_type=owner_type_snapshot,
|
||||
owner_id=owner_id_snapshot,
|
||||
for_update=False,
|
||||
)
|
||||
if owner is None or not is_attempt_current(owner, attempt_snapshot):
|
||||
return None
|
||||
|
||||
priority = int(
|
||||
settings.DOWNLOAD_TASK_PRIORITY_RECOVER
|
||||
@@ -325,7 +355,7 @@ async def enqueue_download_task(
|
||||
detail={"reason": reason, "error": str(exc)},
|
||||
)
|
||||
return None
|
||||
return owner.download_celery_task_id
|
||||
return str(owner.download_celery_task_id or celery_task_id_snapshot or "") or None
|
||||
|
||||
|
||||
async def _claim(
|
||||
@@ -333,9 +363,9 @@ async def _claim(
|
||||
owner: GenerationOwner,
|
||||
*,
|
||||
claim_token: str,
|
||||
) -> bool:
|
||||
) -> GenerationOwner | None:
|
||||
if not owner_is_generating(owner) or owner_is_completed(owner):
|
||||
return False
|
||||
return None
|
||||
allowed = {
|
||||
_stage(owner, ChatGenerationPipelineStage.RESULT_READY),
|
||||
_stage(owner, ChatGenerationPipelineStage.DOWNLOAD_QUEUED),
|
||||
@@ -343,7 +373,7 @@ async def _claim(
|
||||
_stage(owner, ChatGenerationPipelineStage.RETRY_WAITING),
|
||||
}
|
||||
if owner.pipeline_stage not in allowed:
|
||||
return False
|
||||
return None
|
||||
|
||||
now = _now()
|
||||
# Redis execution lock is authoritative. A database lease left by a
|
||||
@@ -356,7 +386,7 @@ async def _claim(
|
||||
and next_retry
|
||||
and next_retry > now
|
||||
):
|
||||
return False
|
||||
return None
|
||||
|
||||
owner.pipeline_stage = _stage(
|
||||
owner, ChatGenerationPipelineStage.DOWNLOADING
|
||||
@@ -366,7 +396,18 @@ async def _claim(
|
||||
owner.download_lease_until = _lease_until()
|
||||
owner.download_attempt_count = int(owner.download_attempt_count or 0) + 1
|
||||
owner.download_last_error = None
|
||||
owner_type_snapshot = owner_type_of(owner)
|
||||
owner_id_snapshot = str(owner.id)
|
||||
attempt_snapshot = int(owner.generation_attempt_no or 1)
|
||||
await db.commit()
|
||||
owner = await _reload_owner_after_commit(
|
||||
db,
|
||||
owner_type=owner_type_snapshot,
|
||||
owner_id=owner_id_snapshot,
|
||||
attempt_no=attempt_snapshot,
|
||||
)
|
||||
if owner is None:
|
||||
return None
|
||||
await _register_active(
|
||||
owner,
|
||||
check_at=owner.download_lease_until,
|
||||
@@ -382,7 +423,7 @@ async def _claim(
|
||||
"claim_token_suffix": claim_token[-8:],
|
||||
},
|
||||
)
|
||||
return True
|
||||
return owner
|
||||
|
||||
|
||||
async def _sync_snapshot(db: AsyncSession, owner: GenerationOwner) -> None:
|
||||
@@ -446,12 +487,33 @@ async def _mark_failed(
|
||||
owner.download_last_error = error_message
|
||||
owner.download_lease_until = None
|
||||
owner.download_next_retry_at = None
|
||||
owner_type_snapshot = owner_type_of(owner)
|
||||
owner_id_snapshot = str(owner.id)
|
||||
attempt_snapshot = int(owner.generation_attempt_no or 1)
|
||||
generation_mode_snapshot = str(getattr(owner, "generation_mode", "") or "") or None
|
||||
await db.commit()
|
||||
await notify_owner_finished(db, owner)
|
||||
owner = await _reload_owner_after_commit(
|
||||
db,
|
||||
owner_type=owner_type_snapshot,
|
||||
owner_id=owner_id_snapshot,
|
||||
attempt_no=attempt_snapshot,
|
||||
)
|
||||
if owner is not None:
|
||||
await notify_owner_finished(db, owner)
|
||||
await db.commit()
|
||||
await _remove_active(owner)
|
||||
owner = await _reload_owner_after_commit(
|
||||
db,
|
||||
owner_type=owner_type_snapshot,
|
||||
owner_id=owner_id_snapshot,
|
||||
attempt_no=attempt_snapshot,
|
||||
)
|
||||
if owner is not None:
|
||||
await _remove_active(owner)
|
||||
await log_task_event(
|
||||
owner,
|
||||
owner_type=owner_type_snapshot,
|
||||
owner_id=owner_id_snapshot,
|
||||
generation_attempt_no=attempt_snapshot,
|
||||
generation_mode=generation_mode_snapshot,
|
||||
event_type=(
|
||||
ChatGenerationTaskEventType.DOWNLOAD_FAILED_NON_RETRYABLE.value
|
||||
if non_retryable
|
||||
@@ -481,7 +543,18 @@ async def _schedule_retry(
|
||||
owner.download_celery_task_id = _build_celery_task_id(
|
||||
owner, reason="retry"
|
||||
)
|
||||
owner_type_snapshot = owner_type_of(owner)
|
||||
owner_id_snapshot = str(owner.id)
|
||||
attempt_snapshot = int(owner.generation_attempt_no or 1)
|
||||
await db.commit()
|
||||
owner = await _reload_owner_after_commit(
|
||||
db,
|
||||
owner_type=owner_type_snapshot,
|
||||
owner_id=owner_id_snapshot,
|
||||
attempt_no=attempt_snapshot,
|
||||
)
|
||||
if owner is None:
|
||||
return
|
||||
await _register_active(
|
||||
owner,
|
||||
check_at=owner.download_next_retry_at,
|
||||
@@ -555,7 +628,18 @@ async def _restore_after_lock_error(
|
||||
0, int(owner.download_attempt_count or 0) - 1
|
||||
)
|
||||
owner.download_enqueued_at = _now()
|
||||
owner_type_snapshot = owner_type_of(owner)
|
||||
owner_id_snapshot = str(owner.id)
|
||||
attempt_snapshot = int(owner.generation_attempt_no or 1)
|
||||
await db.commit()
|
||||
owner = await _reload_owner_after_commit(
|
||||
db,
|
||||
owner_type=owner_type_snapshot,
|
||||
owner_id=owner_id_snapshot,
|
||||
attempt_no=attempt_snapshot,
|
||||
)
|
||||
if owner is None:
|
||||
return
|
||||
await _register_active(
|
||||
owner,
|
||||
check_at=_queue_timeout_at(),
|
||||
@@ -631,7 +715,8 @@ async def _run(
|
||||
if not is_attempt_current(owner, effective_attempt):
|
||||
await _remove_active(owner)
|
||||
return
|
||||
if not await _claim(db, owner, claim_token=lease.token):
|
||||
owner = await _claim(db, owner, claim_token=lease.token)
|
||||
if owner is None:
|
||||
return
|
||||
claimed = True
|
||||
|
||||
@@ -705,15 +790,26 @@ async def _run(
|
||||
owner.download_lease_until = None
|
||||
owner.download_next_retry_at = None
|
||||
owner.download_last_error = None
|
||||
await db.commit()
|
||||
await _remove_active(owner)
|
||||
upscale_mode = str(getattr(owner, "generation_mode", "") or "") or None
|
||||
upscale_stage = str(owner.pipeline_stage or "") or None
|
||||
await enqueue_upscale_task(
|
||||
db, upscale=upscale, reason="source_download_completed"
|
||||
)
|
||||
owner = await _reload_owner_after_commit(
|
||||
db,
|
||||
owner_type=normalized_owner_type,
|
||||
owner_id=task_id,
|
||||
attempt_no=effective_attempt,
|
||||
)
|
||||
if owner is not None:
|
||||
await _remove_active(owner)
|
||||
await log_task_event(
|
||||
owner,
|
||||
owner_type=normalized_owner_type,
|
||||
owner_id=task_id,
|
||||
generation_attempt_no=effective_attempt,
|
||||
generation_mode=upscale_mode,
|
||||
event_type=ChatGenerationTaskEventType.DOWNLOAD_SUCCESS.value,
|
||||
to_stage=owner.pipeline_stage,
|
||||
to_stage=(str(owner.pipeline_stage or "") if owner is not None else upscale_stage),
|
||||
detail={
|
||||
"upscale_source_path": downloaded.storage_path
|
||||
},
|
||||
@@ -739,14 +835,33 @@ async def _run(
|
||||
owner.download_last_error = None
|
||||
await _record_resource(db, owner, downloaded)
|
||||
await _sync_snapshot(db, owner)
|
||||
completion_mode = str(getattr(owner, "generation_mode", "") or "") or None
|
||||
completion_stage = str(owner.pipeline_stage or "") or None
|
||||
await db.commit()
|
||||
await notify_owner_finished(db, owner)
|
||||
owner = await _reload_owner_after_commit(
|
||||
db,
|
||||
owner_type=normalized_owner_type,
|
||||
owner_id=task_id,
|
||||
attempt_no=effective_attempt,
|
||||
)
|
||||
if owner is not None:
|
||||
await notify_owner_finished(db, owner)
|
||||
await db.commit()
|
||||
await _remove_active(owner)
|
||||
owner = await _reload_owner_after_commit(
|
||||
db,
|
||||
owner_type=normalized_owner_type,
|
||||
owner_id=task_id,
|
||||
attempt_no=effective_attempt,
|
||||
)
|
||||
if owner is not None:
|
||||
await _remove_active(owner)
|
||||
await log_task_event(
|
||||
owner,
|
||||
owner_type=normalized_owner_type,
|
||||
owner_id=task_id,
|
||||
generation_attempt_no=effective_attempt,
|
||||
generation_mode=completion_mode,
|
||||
event_type=ChatGenerationTaskEventType.DOWNLOAD_SUCCESS.value,
|
||||
to_stage=owner.pipeline_stage,
|
||||
to_stage=(str(owner.pipeline_stage or "") if owner is not None else completion_stage),
|
||||
detail={
|
||||
"resource_url": downloaded.url,
|
||||
"file_size_bytes": downloaded.file_size_bytes,
|
||||
|
||||
@@ -22,7 +22,7 @@ from app.enums.generation_task import (
|
||||
from app.models.base import async_session
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
from app.services.error_codes import extract_error_message
|
||||
from app.services.generation.log_service import log_provider_call, log_task_event
|
||||
from app.services.generation.log_service import log_task_event
|
||||
from app.services.generation.pipeline.db_lock_service import DatabaseRowLockBusy
|
||||
from app.services.generation.pipeline.lifecycle_service import (
|
||||
mark_owner_failed_and_refund_once,
|
||||
@@ -101,24 +101,6 @@ def _engine_snapshot(owner: GenerationOwner) -> dict:
|
||||
return {}
|
||||
|
||||
|
||||
async def _log_poll_provider_call_after_commit(
|
||||
owner: GenerationOwner,
|
||||
*,
|
||||
provider_response: Any,
|
||||
) -> None:
|
||||
"""Provider logs use an independent session, so the owner row must be committed first."""
|
||||
snapshot = _engine_snapshot(owner)
|
||||
await log_provider_call(
|
||||
owner,
|
||||
provider=snapshot.get("provider") or "ark",
|
||||
api_type=f"{owner.gen_type}_poll",
|
||||
model=snapshot.get("model_name"),
|
||||
engine_id=owner.engine_id,
|
||||
status="success",
|
||||
provider_task_id=owner_provider_task_id(owner),
|
||||
response_data=provider_response,
|
||||
)
|
||||
|
||||
|
||||
def _registry_id(owner: GenerationOwner) -> str:
|
||||
return redis_owner_item_id(
|
||||
@@ -235,6 +217,24 @@ async def remove_poll_active(
|
||||
)
|
||||
|
||||
|
||||
async def _reload_owner_after_commit(
|
||||
db,
|
||||
*,
|
||||
owner_type: str,
|
||||
owner_id: str,
|
||||
attempt_no: int,
|
||||
) -> GenerationOwner | None:
|
||||
fresh = await load_generation_owner(
|
||||
db,
|
||||
owner_type=owner_type,
|
||||
owner_id=owner_id,
|
||||
for_update=False,
|
||||
)
|
||||
if fresh is None or not is_attempt_current(fresh, attempt_no):
|
||||
return None
|
||||
return fresh
|
||||
|
||||
|
||||
async def _sync_snapshot(
|
||||
db, owner: GenerationOwner, provider_response: Any = None
|
||||
) -> None:
|
||||
@@ -266,16 +266,34 @@ async def _mark_failed(
|
||||
owner.next_poll_at = None
|
||||
owner.poll_claim_token = None
|
||||
owner.poll_lease_until = None
|
||||
owner_type_snapshot = owner_type_of(owner)
|
||||
owner_id_snapshot = str(owner.id)
|
||||
attempt_snapshot = int(owner.generation_attempt_no or 1)
|
||||
mode_snapshot = owner_mode(owner)
|
||||
stage_snapshot = str(owner.pipeline_stage or "")
|
||||
await db.commit()
|
||||
await notify_owner_finished(db, owner)
|
||||
await db.commit()
|
||||
await remove_poll_active(owner)
|
||||
fresh_owner = await _reload_owner_after_commit(
|
||||
db,
|
||||
owner_type=owner_type_snapshot,
|
||||
owner_id=owner_id_snapshot,
|
||||
attempt_no=attempt_snapshot,
|
||||
)
|
||||
if fresh_owner is not None:
|
||||
await notify_owner_finished(db, fresh_owner)
|
||||
await remove_poll_active(
|
||||
owner_type=owner_type_snapshot,
|
||||
owner_id=owner_id_snapshot,
|
||||
attempt_no=attempt_snapshot,
|
||||
)
|
||||
await log_task_event(
|
||||
owner,
|
||||
owner_type=owner_type_snapshot,
|
||||
owner_id=owner_id_snapshot,
|
||||
generation_attempt_no=attempt_snapshot,
|
||||
generation_mode=mode_snapshot,
|
||||
event_type=event_type,
|
||||
message=message,
|
||||
detail=detail,
|
||||
to_stage=owner.pipeline_stage,
|
||||
to_stage=stage_snapshot,
|
||||
)
|
||||
|
||||
|
||||
@@ -302,15 +320,30 @@ async def _schedule_next_poll(
|
||||
owner.poll_interval_seconds = schedule.poll_interval_seconds
|
||||
owner.poll_claim_token = None
|
||||
owner.poll_lease_until = None
|
||||
owner_type_snapshot = owner_type_of(owner)
|
||||
owner_id_snapshot = str(owner.id)
|
||||
attempt_snapshot = int(owner.generation_attempt_no or 1)
|
||||
mode_snapshot = owner_mode(owner)
|
||||
await db.commit()
|
||||
fresh_owner = await _reload_owner_after_commit(
|
||||
db,
|
||||
owner_type=owner_type_snapshot,
|
||||
owner_id=owner_id_snapshot,
|
||||
attempt_no=attempt_snapshot,
|
||||
)
|
||||
if fresh_owner is None:
|
||||
return
|
||||
await register_poll_active(
|
||||
owner,
|
||||
fresh_owner,
|
||||
check_at=schedule.next_poll_at,
|
||||
next_poll_at=schedule.next_poll_at,
|
||||
reason=schedule.reason,
|
||||
)
|
||||
await log_task_event(
|
||||
owner,
|
||||
owner_type=owner_type_snapshot,
|
||||
owner_id=owner_id_snapshot,
|
||||
generation_attempt_no=attempt_snapshot,
|
||||
generation_mode=mode_snapshot,
|
||||
event_type=ChatGenerationTaskEventType.POLL_SCHEDULED.value,
|
||||
message=f"已登记下一次轮询。reason={schedule.reason}",
|
||||
detail={
|
||||
@@ -320,10 +353,10 @@ async def _schedule_next_poll(
|
||||
)
|
||||
if schedule.direct_countdown:
|
||||
poll_generation_task.apply_async(
|
||||
args=[owner.id],
|
||||
args=[owner_id_snapshot],
|
||||
kwargs={
|
||||
"owner_type": owner_type_of(owner),
|
||||
"generation_attempt_no": int(owner.generation_attempt_no or 1),
|
||||
"owner_type": owner_type_snapshot,
|
||||
"generation_attempt_no": attempt_snapshot,
|
||||
"force_due": False,
|
||||
},
|
||||
queue=POLL_QUEUE,
|
||||
@@ -381,13 +414,21 @@ async def _restore_after_lock_error(
|
||||
owner.poll_claim_token = None
|
||||
owner.poll_lease_until = None
|
||||
owner.next_poll_at = _now()
|
||||
next_poll_at = owner.next_poll_at
|
||||
await db.commit()
|
||||
await register_poll_active(
|
||||
owner,
|
||||
check_at=owner.next_poll_at,
|
||||
next_poll_at=owner.next_poll_at,
|
||||
reason="poll_execution_lock_error",
|
||||
fresh_owner = await _reload_owner_after_commit(
|
||||
db,
|
||||
owner_type=owner_type,
|
||||
owner_id=owner_id,
|
||||
attempt_no=attempt_no,
|
||||
)
|
||||
if fresh_owner is not None:
|
||||
await register_poll_active(
|
||||
fresh_owner,
|
||||
check_at=next_poll_at,
|
||||
next_poll_at=next_poll_at,
|
||||
reason="poll_execution_lock_error",
|
||||
)
|
||||
|
||||
|
||||
async def _run(
|
||||
@@ -478,13 +519,21 @@ async def _run(
|
||||
# A stale database lease left by a crashed worker must not block the
|
||||
# worker that successfully acquired the current Redis lock.
|
||||
if not force_due and is_poll_not_due(owner, now=current):
|
||||
next_poll_at = owner.next_poll_at
|
||||
await db.commit()
|
||||
await register_poll_active(
|
||||
owner,
|
||||
check_at=owner.next_poll_at,
|
||||
next_poll_at=owner.next_poll_at,
|
||||
reason="poll_task_not_due",
|
||||
owner = await _reload_owner_after_commit(
|
||||
db,
|
||||
owner_type=normalized_owner_type,
|
||||
owner_id=task_id,
|
||||
attempt_no=effective_attempt,
|
||||
)
|
||||
if owner is not None:
|
||||
await register_poll_active(
|
||||
owner,
|
||||
check_at=next_poll_at,
|
||||
next_poll_at=next_poll_at,
|
||||
reason="poll_task_not_due",
|
||||
)
|
||||
return
|
||||
|
||||
final_poll = is_final_poll_due(owner, now=current)
|
||||
@@ -507,12 +556,22 @@ async def _run(
|
||||
owner.poll_lease_until = _poll_lease_until(current)
|
||||
owner.poll_count = int(owner.poll_count or 0) + 1
|
||||
owner.last_poll_at = current
|
||||
poll_lease_until = owner.poll_lease_until
|
||||
next_poll_at = owner.next_poll_at
|
||||
await db.commit()
|
||||
claim_started = True
|
||||
owner = await _reload_owner_after_commit(
|
||||
db,
|
||||
owner_type=normalized_owner_type,
|
||||
owner_id=task_id,
|
||||
attempt_no=effective_attempt,
|
||||
)
|
||||
if owner is None:
|
||||
return
|
||||
await register_poll_active(
|
||||
owner,
|
||||
check_at=owner.poll_lease_until,
|
||||
next_poll_at=owner.next_poll_at,
|
||||
check_at=poll_lease_until,
|
||||
next_poll_at=next_poll_at,
|
||||
reason="polling_lease",
|
||||
)
|
||||
|
||||
@@ -562,9 +621,6 @@ async def _run(
|
||||
event_type=ChatGenerationTaskEventType.POLL_FAILED.value,
|
||||
detail=poll_result,
|
||||
)
|
||||
await _log_poll_provider_call_after_commit(
|
||||
owner, provider_response=provider_response
|
||||
)
|
||||
return
|
||||
owner.pipeline_stage = _stage(
|
||||
owner, ChatGenerationPipelineStage.RESULT_READY
|
||||
@@ -573,16 +629,30 @@ async def _run(
|
||||
owner.poll_claim_token = None
|
||||
owner.poll_lease_until = None
|
||||
owner.next_poll_at = None
|
||||
success_stage = str(owner.pipeline_stage or "")
|
||||
success_mode = owner_mode(owner)
|
||||
await db.commit()
|
||||
await _log_poll_provider_call_after_commit(
|
||||
owner, provider_response=provider_response
|
||||
owner = await _reload_owner_after_commit(
|
||||
db,
|
||||
owner_type=normalized_owner_type,
|
||||
owner_id=task_id,
|
||||
attempt_no=effective_attempt,
|
||||
)
|
||||
await remove_poll_active(
|
||||
owner_type=normalized_owner_type,
|
||||
owner_id=task_id,
|
||||
attempt_no=effective_attempt,
|
||||
)
|
||||
await remove_poll_active(owner)
|
||||
await log_task_event(
|
||||
owner,
|
||||
owner_type=normalized_owner_type,
|
||||
owner_id=task_id,
|
||||
generation_attempt_no=effective_attempt,
|
||||
generation_mode=success_mode,
|
||||
event_type=ChatGenerationTaskEventType.POLL_SUCCESS.value,
|
||||
to_stage=owner.pipeline_stage,
|
||||
to_stage=success_stage,
|
||||
)
|
||||
if owner is None:
|
||||
return
|
||||
from app.tasks.generation_download_tasks import (
|
||||
enqueue_download_task,
|
||||
)
|
||||
@@ -603,9 +673,6 @@ async def _run(
|
||||
event_type=ChatGenerationTaskEventType.POLL_FAILED.value,
|
||||
detail=poll_result,
|
||||
)
|
||||
await _log_poll_provider_call_after_commit(
|
||||
owner, provider_response=provider_response
|
||||
)
|
||||
return
|
||||
|
||||
if final_poll:
|
||||
@@ -617,18 +684,12 @@ async def _run(
|
||||
event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value,
|
||||
detail=poll_result,
|
||||
)
|
||||
await _log_poll_provider_call_after_commit(
|
||||
owner, provider_response=provider_response
|
||||
)
|
||||
return
|
||||
|
||||
owner.poll_error_count = 0
|
||||
await _schedule_next_poll(
|
||||
db, owner, reason="poll_pending_next"
|
||||
)
|
||||
await _log_poll_provider_call_after_commit(
|
||||
owner, provider_response=provider_response
|
||||
)
|
||||
await log_task_event(
|
||||
owner,
|
||||
event_type=ChatGenerationTaskEventType.POLL_PENDING.value,
|
||||
|
||||
@@ -27,12 +27,97 @@ async def _recover_generation_records_once(*, include_create: bool, include_poll
|
||||
from app.tasks.generation_poll_tasks import poll_generation_task
|
||||
from app.tasks.generation_download_tasks import enqueue_download_task
|
||||
|
||||
counts: dict[str, Any] = {"create": 0, "poll": 0, "download": 0, "errors": []}
|
||||
counts: dict[str, Any] = {
|
||||
"create": 0,
|
||||
"poll": 0,
|
||||
"download": 0,
|
||||
"inconsistent": 0,
|
||||
"inconsistent_timeout": 0,
|
||||
"errors": [],
|
||||
}
|
||||
batch_size = max(1, int(settings.GENERATION_RECOVERY_BATCH_SIZE or 20))
|
||||
cursor = None
|
||||
async with async_session() as db:
|
||||
while True:
|
||||
batch = await find_generation_record_recovery_batch(db, limit=batch_size, cursor=cursor)
|
||||
if include_create or include_poll:
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from app.enums.generation_status import GenerationRecordPipelineStage
|
||||
from app.enums.generation_task import ChatGenerationTaskEventType, GenerationMode
|
||||
from app.services.generation.log_service import log_task_event
|
||||
from app.services.generation.pipeline.owner_service import load_generation_owner
|
||||
from app.services.generation.refund_service import mark_generation_record_failed_and_refund_once
|
||||
from app.services.redis_registry_service import ensure_aware_utc
|
||||
|
||||
for ref in batch.inconsistent:
|
||||
try:
|
||||
owner = await load_generation_owner(
|
||||
db,
|
||||
owner_type=ref.owner_type,
|
||||
owner_id=ref.owner_id,
|
||||
for_update=True,
|
||||
)
|
||||
if (
|
||||
owner is None
|
||||
or int(owner.generation_attempt_no or 1)
|
||||
!= int(ref.generation_attempt_no or 1)
|
||||
):
|
||||
await db.rollback()
|
||||
continue
|
||||
deadline_at = ensure_aware_utc(getattr(owner, "deadline_at", None))
|
||||
owner_id_snapshot = str(owner.id)
|
||||
attempt_snapshot = int(owner.generation_attempt_no or 1)
|
||||
previous_stage = str(owner.pipeline_stage or "")
|
||||
if deadline_at is not None and deadline_at <= datetime.now(timezone.utc):
|
||||
owner.pipeline_stage = GenerationRecordPipelineStage.TIMEOUT.value
|
||||
await mark_generation_record_failed_and_refund_once(
|
||||
db,
|
||||
record=owner,
|
||||
generation_attempt_no=ref.generation_attempt_no,
|
||||
error_message="恢复证据异常且已超过任务截止时间",
|
||||
)
|
||||
await db.commit()
|
||||
await log_task_event(
|
||||
owner_type=ref.owner_type,
|
||||
owner_id=owner_id_snapshot,
|
||||
generation_attempt_no=attempt_snapshot,
|
||||
generation_mode=GenerationMode.GENERATION_RECORD.value,
|
||||
event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value,
|
||||
from_stage=previous_stage,
|
||||
to_stage=GenerationRecordPipelineStage.TIMEOUT.value,
|
||||
message="恢复证据异常任务超过截止时间,已失败并幂等退款",
|
||||
)
|
||||
counts["inconsistent_timeout"] += 1
|
||||
continue
|
||||
if previous_stage != GenerationRecordPipelineStage.RECOVERY_INCONSISTENT.value:
|
||||
owner.pipeline_stage = GenerationRecordPipelineStage.RECOVERY_INCONSISTENT.value
|
||||
owner.error_message = (
|
||||
f"恢复证据异常:阶段 {previous_stage} 缺少 remote_result_url 和供应商任务ID"
|
||||
)
|
||||
await db.commit()
|
||||
await log_task_event(
|
||||
owner_type=ref.owner_type,
|
||||
owner_id=owner_id_snapshot,
|
||||
generation_attempt_no=attempt_snapshot,
|
||||
generation_mode=GenerationMode.GENERATION_RECORD.value,
|
||||
event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_INCONSISTENT.value,
|
||||
from_stage=previous_stage,
|
||||
to_stage=GenerationRecordPipelineStage.RECOVERY_INCONSISTENT.value,
|
||||
message="恢复证据异常,已隔离且不重新创建供应商任务",
|
||||
)
|
||||
else:
|
||||
await db.rollback()
|
||||
counts["inconsistent"] += 1
|
||||
except Exception as exc:
|
||||
await db.rollback()
|
||||
counts["errors"].append(
|
||||
{
|
||||
"owner_id": ref.owner_id,
|
||||
"stage": "recovery_inconsistent",
|
||||
"error": str(exc),
|
||||
}
|
||||
)
|
||||
if include_create:
|
||||
for ref in batch.create:
|
||||
try:
|
||||
|
||||
Reference in New Issue
Block a user