修复行业智造图片生成BUG

This commit is contained in:
2026-07-16 16:54:42 +08:00
parent a7c99839de
commit 395702fbd5
+97 -50
View File
@@ -3,6 +3,7 @@ import json
import logging import logging
import os import os
from datetime import datetime, timezone from datetime import datetime, timezone
from urllib.parse import urlparse
from sqlalchemy import or_, select from sqlalchemy import or_, select
@@ -49,6 +50,31 @@ def _source_date_dir(record: GenerationRecord) -> str:
return created.strftime("%Y/%m/%d") return created.strftime("%Y/%m/%d")
def _normalize_image_extension(output_format: str | None, remote_url: str | None = None) -> str:
value = str(output_format or "").strip().lower()
if value in {"jpg", "jpeg"}:
return "jpg"
if value == "png":
return "png"
if value == "webp":
return "webp"
if remote_url:
try:
path = urlparse(remote_url).path or ""
except Exception:
path = str(remote_url)
suffix = os.path.splitext(path)[1].lower().lstrip(".")
if suffix in {"jpg", "jpeg"}:
return "jpg"
if suffix == "png":
return "png"
if suffix == "webp":
return "webp"
return "jpg"
async def _download_generation_record_upscale_source( async def _download_generation_record_upscale_source(
record: GenerationRecord, record: GenerationRecord,
remote_url: str, remote_url: str,
@@ -308,10 +334,10 @@ class TaskQueue:
await asyncio.sleep(POLL_INTERVAL) await asyncio.sleep(POLL_INTERVAL)
await self.queue.put(record_id) await self.queue.put(record_id)
async def _process_image(self, db, record): async def _process_image(self, db, record: GenerationRecord):
"""Process image generation task - calls API directly.""" """Process image generation task - calls API directly."""
record_id = record.id record_id = record.id
from app.services.image_gen import submit_image_task, _log_image_response from app.services.image_gen import submit_image_task
try: try:
engine = await get_active_image_engine(db) engine = await get_active_image_engine(db)
@@ -323,60 +349,81 @@ class TaskQueue:
include_media_references=False, include_media_references=False,
) )
if poll_result["error"] == "": if not isinstance(poll_result, dict):
remote_url = poll_result.get("image_url") raise RuntimeError("图片供应商返回结构异常")
storage_path = None
file_size_bytes = 0
if settings.STORAGE_TYPE == "local" and remote_url:
try:
date_dir = datetime.now().strftime("%Y/%m/%d")
dest_dir = os.path.join(settings.STORAGE_IMAGE_LOCAL_PATH, date_dir)
os.makedirs(dest_dir, exist_ok=True)
dest = os.path.join(dest_dir, f"{record_id}.png")
await download_image(remote_url, dest)
record.image_url = f"/generate/images/{date_dir}/{record_id}.png"
storage_path = dest
file_size_bytes = safe_file_size(dest)
except Exception as e:
logger.warning(f"Download failed, using remote URL: {e}")
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:
await record_generation_record_generated_resource(
db,
record,
resource_url=record.image_url,
storage_path=storage_path,
file_size_bytes=file_size_bytes,
remote_url=remote_url,
generated_at=record.generated_at,
)
await db.commit()
logger.info(f"Image task completed: {record_id}")
else:
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message=poll_result.get("error", "图片生成失败"),
)
await db.commit()
logger.info(f"Image task failed: {record_id}")
_log_image_response(record_id, poll_result)
except Exception as e: items = poll_result.get("items") or []
if not isinstance(items, list):
raise RuntimeError("图片供应商返回结果列表异常")
if not items:
raise RuntimeError("图片供应商未返回图片结果")
if len(items) != 1:
raise RuntimeError(f"图片供应商单图返回数量异常,期望 1,实际 {len(items)}")
item = items[0] or {}
if not isinstance(item, dict):
raise RuntimeError("图片供应商返回单项结果结构异常")
item_error = item.get("error_message") or item.get("error_code")
if item_error:
raise RuntimeError(str(item_error))
remote_url = str(item.get("remote_result_url") or "").strip()
if not remote_url:
raise RuntimeError("图片供应商成功响应但没有图片地址")
storage_path = None
file_size_bytes = 0
if settings.STORAGE_TYPE == "local":
try:
date_dir = _source_date_dir(record)
dest_dir = os.path.join(settings.STORAGE_IMAGE_LOCAL_PATH, date_dir)
os.makedirs(dest_dir, exist_ok=True)
extension = _normalize_image_extension(item.get("output_format"), remote_url)
dest = os.path.join(dest_dir, f"{record_id}.{extension}")
await download_image(remote_url, dest)
record.image_url = f"/generate/images/{date_dir}/{record_id}.{extension}"
storage_path = dest
file_size_bytes = safe_file_size(dest)
except Exception as exc:
logger.warning("GenerationRecord 图片本地保存失败,回退远程地址: record_id=%s error=%s", record_id, exc)
record.image_url = remote_url
else:
record.image_url = remote_url
record.image_tokens_used = int(poll_result.get("image_tokens", 0) or 0)
provider_response = poll_result.get("response_data") or {}
await sync_generation_record_media_token_snapshot(
db,
record,
provider_response=provider_response if isinstance(provider_response, dict) else {},
)
record.status = "completed"
record.pipeline_stage = GenerationRecordPipelineStage.DONE.value
record.generated_at = datetime.now(timezone.utc)
record.error_message = None
if record.image_url:
await record_generation_record_generated_resource(
db,
record,
resource_url=record.image_url,
storage_path=storage_path,
file_size_bytes=file_size_bytes,
remote_url=remote_url,
generated_at=record.generated_at,
)
await db.commit()
logger.info("Image task completed: %s", record_id)
except Exception as exc:
record.pipeline_stage = GenerationRecordPipelineStage.FAILED.value
await mark_generation_record_failed_and_refund_once( await mark_generation_record_failed_and_refund_once(
db, db,
record=record, record=record,
error_message=str(e), error_message=(getattr(exc, "safe_message", None) or str(exc) or "图片生成失败"),
) )
_log_image_response(record_id, {}, str(e))
await db.commit() await db.commit()
logger.error(f"Image task failed: {record_id}, error: {e}") logger.error("Image task failed: %s, error: %s", record_id, exc, exc_info=True)
def stop(self): def stop(self):
"""Signal the queue to stop.""" """Signal the queue to stop."""