1
This commit is contained in:
@@ -3,14 +3,21 @@ import json
|
||||
import logging
|
||||
import os
|
||||
from datetime import datetime
|
||||
from types import SimpleNamespace
|
||||
|
||||
from sqlalchemy import select
|
||||
|
||||
from app.models.base import async_session
|
||||
from app.models.generation_record import GenerationRecord
|
||||
from app.services.video_gen import get_active_engine, poll_task_status, download_video, _log_video_response
|
||||
from app.services.image_gen import get_active_image_engine, download_image
|
||||
from app.services.media_token_usage_snapshot_service import sync_generation_record_media_token_snapshot
|
||||
from app.models.image_engine import ImageEngine
|
||||
from app.models.video_engine import VideoEngine
|
||||
from app.enums.model_pricing import PricingSnapshotStage, ProviderCostStatus
|
||||
from app.services.video_gen import poll_task_status, download_video, _log_video_response
|
||||
from app.services.image_gen import download_image, is_sync_image_provider_result_uncertain
|
||||
from app.services.media_token_usage_snapshot_service import (
|
||||
mark_media_provider_cost_status,
|
||||
sync_generation_record_media_token_snapshot,
|
||||
)
|
||||
from app.services.resource_accounting_service import (
|
||||
record_generation_record_generated_resource,
|
||||
safe_file_size,
|
||||
@@ -25,6 +32,38 @@ POLL_INTERVAL = 30 # seconds between polls
|
||||
MAX_POLLS = 60 # max 30 minutes total
|
||||
|
||||
|
||||
def _loads(data: str | None) -> dict:
|
||||
if not data:
|
||||
return {}
|
||||
try:
|
||||
value = json.loads(data)
|
||||
return value if isinstance(value, dict) else {}
|
||||
except Exception:
|
||||
return {}
|
||||
|
||||
|
||||
async def _get_runtime_engine(db, record: GenerationRecord):
|
||||
"""使用生成开始时冻结的引擎快照,数据库行只读取当前密钥。"""
|
||||
if not record.engine_id:
|
||||
raise ValueError("生成记录缺少锁定的 engine_id")
|
||||
snapshot = _loads(record.engine_snapshot_json)
|
||||
model = ImageEngine if record.gen_type == "image" else VideoEngine
|
||||
engine = (await db.execute(select(model).where(model.id == record.engine_id).limit(1))).scalar_one_or_none()
|
||||
if not engine:
|
||||
raise ValueError("锁定的生成引擎不存在")
|
||||
return SimpleNamespace(
|
||||
id=record.engine_id,
|
||||
name=snapshot.get("name") or engine.name,
|
||||
provider=snapshot.get("provider") or engine.provider,
|
||||
api_base=snapshot.get("api_base") or engine.api_base,
|
||||
api_key=engine.api_key,
|
||||
model_name=snapshot.get("model_name") or engine.model_name,
|
||||
generate_url=snapshot.get("generate_url") or getattr(engine, "generate_url", ""),
|
||||
query_url=snapshot.get("query_url") or getattr(engine, "query_url", ""),
|
||||
default_size=snapshot.get("default_size") or getattr(engine, "default_size", "2K"),
|
||||
)
|
||||
|
||||
|
||||
class TaskQueue:
|
||||
def __init__(self):
|
||||
self.queue: asyncio.Queue[str] = asyncio.Queue()
|
||||
@@ -93,7 +132,7 @@ class TaskQueue:
|
||||
async def _process_video(self, db, record):
|
||||
"""Process video generation task."""
|
||||
record_id = record.id
|
||||
|
||||
|
||||
if not record.seedance_task_id:
|
||||
await mark_generation_record_failed_and_refund_once(
|
||||
db,
|
||||
@@ -104,7 +143,7 @@ class TaskQueue:
|
||||
return
|
||||
|
||||
try:
|
||||
engine = await get_active_engine(db)
|
||||
engine = await _get_runtime_engine(db, record)
|
||||
poll_result = await poll_task_status(engine, record.seedance_task_id)
|
||||
except Exception as e:
|
||||
logger.error(f"Poll error for {record_id}: {e}")
|
||||
@@ -133,6 +172,17 @@ class TaskQueue:
|
||||
|
||||
if status == "succeeded":
|
||||
file_url = poll_result.get("video_url", "")
|
||||
resp_data["video_url"] = file_url
|
||||
resp_data["task_id"] = record.seedance_task_id
|
||||
record.provider_response_json = json.dumps(resp_data, ensure_ascii=False, default=str)
|
||||
record.video_tokens_used = poll_result.get("video_tokens", 0)
|
||||
# 视频 Provider 已完成时核算供应商成本;下载失败不影响已发生的供应商费用。
|
||||
await sync_generation_record_media_token_snapshot(
|
||||
db,
|
||||
record,
|
||||
provider_response=resp_data,
|
||||
stage=PricingSnapshotStage.PROVIDER_ASYNC_COMPLETED.value,
|
||||
)
|
||||
storage_path = None
|
||||
file_size_bytes = 0
|
||||
if settings.STORAGE_TYPE == "local" and file_url:
|
||||
@@ -157,8 +207,6 @@ class TaskQueue:
|
||||
record.video_url = file_url
|
||||
else:
|
||||
record.video_url = file_url
|
||||
record.video_tokens_used = poll_result.get("video_tokens", 0)
|
||||
await sync_generation_record_media_token_snapshot(db, record, provider_response=resp_data)
|
||||
record.status = "completed"
|
||||
record.generated_at = datetime.now()
|
||||
if record.video_url:
|
||||
@@ -171,6 +219,12 @@ class TaskQueue:
|
||||
remote_url=file_url,
|
||||
generated_at=record.generated_at,
|
||||
)
|
||||
await sync_generation_record_media_token_snapshot(
|
||||
db,
|
||||
record,
|
||||
provider_response=resp_data,
|
||||
stage=PricingSnapshotStage.RESOURCE_DOWNLOAD_COMPLETED.value,
|
||||
)
|
||||
self._active.pop(record_id, None)
|
||||
await db.commit()
|
||||
logger.info(f"Video task completed: {record_id}")
|
||||
@@ -207,8 +261,11 @@ class TaskQueue:
|
||||
record_id = record.id
|
||||
from app.services.image_gen import submit_image_task, _log_image_response
|
||||
|
||||
provider_call_started = False
|
||||
provider_call_completed = False
|
||||
try:
|
||||
engine = await get_active_image_engine(db)
|
||||
engine = await _get_runtime_engine(db, record)
|
||||
provider_call_started = True
|
||||
poll_result = await asyncio.to_thread(
|
||||
submit_image_task,
|
||||
db,
|
||||
@@ -216,9 +273,23 @@ class TaskQueue:
|
||||
record,
|
||||
include_media_references=False,
|
||||
)
|
||||
provider_call_completed = True
|
||||
|
||||
if poll_result["error"] == "":
|
||||
remote_url = poll_result.get("image_url")
|
||||
try:
|
||||
provider_response = json.loads(poll_result.get("response_data") or "{}")
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
provider_response = {}
|
||||
record.provider_response_json = json.dumps(provider_response, ensure_ascii=False, default=str)
|
||||
record.image_tokens_used = poll_result.get("image_tokens", 0)
|
||||
# 火山图片接口为同步生成:最终响应返回后立即完成供应商成本核算。
|
||||
await sync_generation_record_media_token_snapshot(
|
||||
db,
|
||||
record,
|
||||
provider_response=provider_response,
|
||||
stage=PricingSnapshotStage.PROVIDER_SYNC_COMPLETED.value,
|
||||
)
|
||||
storage_path = None
|
||||
file_size_bytes = 0
|
||||
if settings.STORAGE_TYPE == "local" and remote_url:
|
||||
@@ -236,8 +307,6 @@ class TaskQueue:
|
||||
record.image_url = remote_url
|
||||
else:
|
||||
record.image_url = remote_url
|
||||
record.image_tokens_used = poll_result.get("image_tokens", 0)
|
||||
await sync_generation_record_media_token_snapshot(db, record, provider_response=poll_result)
|
||||
record.status = "completed"
|
||||
record.generated_at = datetime.now()
|
||||
if record.image_url:
|
||||
@@ -250,19 +319,46 @@ class TaskQueue:
|
||||
remote_url=remote_url,
|
||||
generated_at=record.generated_at,
|
||||
)
|
||||
await sync_generation_record_media_token_snapshot(
|
||||
db,
|
||||
record,
|
||||
provider_response=provider_response,
|
||||
stage=PricingSnapshotStage.RESOURCE_DOWNLOAD_COMPLETED.value,
|
||||
)
|
||||
await db.commit()
|
||||
logger.info(f"Image task completed: {record_id}")
|
||||
else:
|
||||
error_message = poll_result.get("error", "图片生成失败")
|
||||
await mark_media_provider_cost_status(
|
||||
db,
|
||||
owner=record,
|
||||
status=ProviderCostStatus.NOT_INCURRED.value,
|
||||
reason=error_message,
|
||||
usage_stage="provider_sync_failed",
|
||||
)
|
||||
await mark_generation_record_failed_and_refund_once(
|
||||
db,
|
||||
record=record,
|
||||
error_message=poll_result.get("error", "图片生成失败"),
|
||||
error_message=error_message,
|
||||
)
|
||||
await db.commit()
|
||||
logger.info(f"Image task failed: {record_id}")
|
||||
_log_image_response(record_id, poll_result)
|
||||
|
||||
except Exception as e:
|
||||
if provider_call_started:
|
||||
uncertain = provider_call_completed or is_sync_image_provider_result_uncertain(e)
|
||||
await mark_media_provider_cost_status(
|
||||
db,
|
||||
owner=record,
|
||||
status=(
|
||||
ProviderCostStatus.PROVIDER_RESULT_UNCERTAIN.value
|
||||
if uncertain
|
||||
else ProviderCostStatus.NOT_INCURRED.value
|
||||
),
|
||||
reason=str(e),
|
||||
usage_stage="provider_sync_exception",
|
||||
)
|
||||
await mark_generation_record_failed_and_refund_once(
|
||||
db,
|
||||
record=record,
|
||||
|
||||
Reference in New Issue
Block a user