修复行业智造图片生成BUG
This commit is contained in:
@@ -3,6 +3,7 @@ import json
|
||||
import logging
|
||||
import os
|
||||
from datetime import datetime, timezone
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from sqlalchemy import or_, select
|
||||
|
||||
@@ -49,6 +50,31 @@ def _source_date_dir(record: GenerationRecord) -> str:
|
||||
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(
|
||||
record: GenerationRecord,
|
||||
remote_url: str,
|
||||
@@ -308,10 +334,10 @@ class TaskQueue:
|
||||
await asyncio.sleep(POLL_INTERVAL)
|
||||
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."""
|
||||
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:
|
||||
engine = await get_active_image_engine(db)
|
||||
@@ -323,60 +349,81 @@ class TaskQueue:
|
||||
include_media_references=False,
|
||||
)
|
||||
|
||||
if poll_result["error"] == "":
|
||||
remote_url = poll_result.get("image_url")
|
||||
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)
|
||||
if not isinstance(poll_result, dict):
|
||||
raise RuntimeError("图片供应商返回结构异常")
|
||||
|
||||
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(
|
||||
db,
|
||||
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()
|
||||
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):
|
||||
"""Signal the queue to stop."""
|
||||
|
||||
Reference in New Issue
Block a user