feat: 添加新的模型,删除后端数据库
This commit is contained in:
+51
-109
@@ -11,16 +11,13 @@ import base64
|
||||
import logging
|
||||
import mimetypes
|
||||
import re
|
||||
from uuid import UUID
|
||||
|
||||
import httpx
|
||||
from sqlalchemy import select
|
||||
|
||||
from config import get_settings
|
||||
from db import get_session_factory
|
||||
from models import Product, ProductAsset, Suite, SuiteImage, SUITE_RUNNING, SUITE_DONE, SUITE_PARTIAL, SUITE_FAILED, STATUS_OK, STATUS_FAILED
|
||||
from services import storage
|
||||
from services.prompt import build_prompt, build_context, type_name, wrap_prompt_for_gpt_edits
|
||||
from services.tasks import Task, TaskImage, TASK_FAILED, TASK_RUNNING, TASK_DONE, TASK_PARTIAL, IMG_OK
|
||||
|
||||
log = logging.getLogger("suite.generator")
|
||||
|
||||
@@ -242,22 +239,25 @@ async def generate_tongyi(prompt: str, ref_images: list[str], size: str = "2048*
|
||||
# 可重试的状态码:中转限流/网关抖动(该中转限流时返回 Cloudflare 502 而非 429)
|
||||
RETRYABLE_STATUS = {429, 500, 502, 503, 504}
|
||||
|
||||
# 中转对 input_fidelity 参数的支持探测:None=未探测,True=支持,False=不支持(已降级)
|
||||
_rightapi_fidelity_supported: bool | None = None
|
||||
# 中转对 input_fidelity 参数的支持探测:按模型记忆不支持该参数的模型(gpt-image 系列支持,
|
||||
# nano-banana 系列可能不认;降级只影响触发过的模型,不牵连其他模型)
|
||||
_rightapi_fidelity_unsupported: set[str] = set()
|
||||
|
||||
|
||||
async def _rightapi_request(s, prompt: str, ref_images: list[str], size: str, model: str) -> bytes:
|
||||
"""gpt-image 系列:有参考图走 /v1/images/edits(multipart),无参考图走 /v1/images/generations。
|
||||
"""RightAPI 各模型:有参考图走 /v1/images/edits(multipart),无参考图走 /v1/images/generations。
|
||||
|
||||
OpenAI 兼容协议:响应固定 b64_json(不支持 response_format 参数,传了报 400);
|
||||
同步调用无任务轮询,高质量档单张 1-5 分钟,超时按文档建议兜底 600s。
|
||||
input_fidelity=high 强制高保真保留输入图细节(商品一致性关键参数,仅 edits 端点);
|
||||
中转若不认该参数(400),自动去掉重试并记住,后续请求不再带。
|
||||
input_fidelity=high 是 gpt-image-1 的 edits 保真参数(gpt-image-2 官方已移除、默认高保真,
|
||||
官逆通道更是不识别);带上是为了兼容按 gpt-image-1 语义实现的中转,中转不认(400)则按模型
|
||||
自动去掉重试并记住,该模型后续请求不再带。
|
||||
"""
|
||||
global _rightapi_fidelity_supported
|
||||
base = s.rightapi_base_url.rstrip("/")
|
||||
headers = {"Authorization": f"Bearer {s.rightapi_api_key}"}
|
||||
use_fidelity = bool(ref_images) and s.rightapi_input_fidelity and _rightapi_fidelity_supported is not False
|
||||
use_fidelity = (
|
||||
bool(ref_images) and s.rightapi_input_fidelity and model not in _rightapi_fidelity_unsupported
|
||||
)
|
||||
|
||||
async with httpx.AsyncClient(timeout=max(s.request_timeout, 600), verify=False) as client:
|
||||
common = {
|
||||
@@ -276,14 +276,12 @@ async def _rightapi_request(s, prompt: str, ref_images: list[str], size: str, mo
|
||||
data, mime = await _resolve_ref_bytes(u)
|
||||
files.append(("image[]", (f"ref-{i + 1}.{mime.split('/')[-1]}", data, mime)))
|
||||
resp = await client.post(f"{base}/v1/images/edits", headers=headers, files=files, data=common)
|
||||
# 中转不认 input_fidelity:去掉参数重试一次(仅一次探测)
|
||||
# 中转不认 input_fidelity:去掉参数重试一次(仅一次探测),降级只记到当前模型
|
||||
if resp.status_code == 400 and use_fidelity and "input_fidelity" in resp.text:
|
||||
_rightapi_fidelity_supported = False
|
||||
log.warning("RightAPI 不支持 input_fidelity 参数,已自动去掉并降级(后续请求不再带)")
|
||||
_rightapi_fidelity_unsupported.add(model)
|
||||
log.warning("RightAPI 模型 %s 不支持 input_fidelity 参数,已自动去掉并降级(该模型后续请求不再带)", model)
|
||||
common.pop("input_fidelity", None)
|
||||
resp = await client.post(f"{base}/v1/images/edits", headers=headers, files=files, data=common)
|
||||
elif resp.is_success and use_fidelity:
|
||||
_rightapi_fidelity_supported = True
|
||||
else:
|
||||
resp = await client.post(
|
||||
f"{base}/v1/images/generations",
|
||||
@@ -365,123 +363,67 @@ def _refs_for_job(images: list[dict], job: dict) -> list[str]:
|
||||
return _order_refs(pool, job.get("kind", ""))
|
||||
|
||||
|
||||
async def _select_ref_images(db, product_id: UUID, type_id: str) -> list[str]:
|
||||
"""商品路径:主图组前几张。转存完成的用本地文件,未完成的直接用源站 URL。"""
|
||||
assets = (await db.scalars(
|
||||
select(ProductAsset).where(
|
||||
ProductAsset.product_id == product_id,
|
||||
ProductAsset.group_key == "main",
|
||||
ProductAsset.type == "img",
|
||||
).order_by(ProductAsset.sort_order)
|
||||
)).all()
|
||||
refs = [a.stored_url or a.source_url for a in assets if (a.stored_url or a.source_url)]
|
||||
if not refs:
|
||||
raise RuntimeError("商品没有可用参考图(未采集主图)")
|
||||
return _order_refs(refs, type_id)
|
||||
# 串行生成队列:所有用户共享同一批 API key,并发生成会触发中转限流
|
||||
# (rightapi 同 key 分钟级冷却);同一时间只跑一个任务,其余保持 pending 排队。
|
||||
_GEN_LOCK = asyncio.Lock()
|
||||
|
||||
|
||||
async def run_suite(suite_id: str) -> None:
|
||||
"""后台执行套图任务:逐张生成 → 落盘 → 记录;单张失败不中断。
|
||||
|
||||
两条路径:
|
||||
- 无状态(product_id 为空):上下文与参考图来自请求自带的 context / ref_images
|
||||
- 商品路径(兼容旧流程):从 product + product_assets 取
|
||||
"""
|
||||
async def run_suite(task: Task) -> None:
|
||||
"""后台执行套图任务:排队 → 逐张生成 → 落盘 → 更新内存状态;单张失败不中断。"""
|
||||
settings = get_settings()
|
||||
async with get_session_factory()() as db:
|
||||
suite = await db.get(Suite, UUID(suite_id))
|
||||
if suite is None:
|
||||
return
|
||||
provider_name = task.provider or settings.image_provider
|
||||
generator = GENERATORS.get(provider_name)
|
||||
if generator is None:
|
||||
task.status = TASK_FAILED
|
||||
task.error = f"未知 provider: {provider_name}"
|
||||
return
|
||||
|
||||
product = None
|
||||
if suite.product_id:
|
||||
product = await db.get(Product, suite.product_id)
|
||||
if product is None:
|
||||
suite.status = SUITE_FAILED
|
||||
suite.error = "商品不存在"
|
||||
await db.commit()
|
||||
return
|
||||
|
||||
suite.status = SUITE_RUNNING
|
||||
await db.commit()
|
||||
|
||||
provider_name = suite.provider or settings.image_provider
|
||||
generator = GENERATORS.get(provider_name)
|
||||
if generator is None:
|
||||
suite.status = SUITE_FAILED
|
||||
suite.error = f"未知 provider: {provider_name}"
|
||||
await db.commit()
|
||||
return
|
||||
|
||||
raw = suite.context if not product else (product.raw or {})
|
||||
ctx = build_context(raw or {}, fallback_name=product.name if product else "")
|
||||
model = suite.model or {
|
||||
"tongyi": settings.dashscope_model,
|
||||
"rightapi": settings.rightapi_image_model,
|
||||
}.get(provider_name, settings.ark_image_model)
|
||||
is_wan = provider_name == "tongyi" and _is_wan_model(model)
|
||||
size = _image_size(provider_name, suite.ratio, is_wan=is_wan, model=model)
|
||||
|
||||
# 任务列表:方案(逐张)优先,旧路径按 types
|
||||
if suite.plan:
|
||||
jobs = [dict(j) for j in suite.plan]
|
||||
else:
|
||||
jobs = [
|
||||
{"kind": t, "title": type_name(t), "detail": "", "prompt_hint": "", "variant_name": None}
|
||||
for t in (suite.types or [])
|
||||
]
|
||||
ctx = build_context(task.context or {}, fallback_name="")
|
||||
model = task.model or {
|
||||
"tongyi": settings.dashscope_model,
|
||||
"rightapi": settings.rightapi_image_model,
|
||||
}.get(provider_name, settings.ark_image_model)
|
||||
is_wan = provider_name == "tongyi" and _is_wan_model(model)
|
||||
size = _image_size(provider_name, task.ratio, is_wan=is_wan, model=model)
|
||||
jobs = [dict(j) for j in task.plan]
|
||||
|
||||
async with _GEN_LOCK:
|
||||
task.status = TASK_RUNNING
|
||||
ok, failed = 0, 0
|
||||
failures: list[str] = []
|
||||
for job in jobs:
|
||||
type_id = job["kind"]
|
||||
image_row = SuiteImage(
|
||||
suite_id=suite.id,
|
||||
type_id=type_id,
|
||||
name=job.get("title") or type_name(type_id),
|
||||
status=STATUS_FAILED,
|
||||
)
|
||||
db.add(image_row)
|
||||
await db.flush()
|
||||
image = TaskImage(type_id=type_id, name=job.get("title") or type_name(type_id))
|
||||
task.images.append(image)
|
||||
try:
|
||||
prompt = build_prompt(
|
||||
type_id, ctx, suite.style_set, suite.lang,
|
||||
extra=job, style_prompt=suite.style_prompt, requirements=suite.requirements,
|
||||
type_id, ctx, task.style_set, task.lang,
|
||||
extra=job, style_prompt=task.style_prompt, requirements=task.requirements,
|
||||
)
|
||||
# gpt-image edits 语义:商品冻结契约前置,防止风格词改商品
|
||||
# gpt-image edits 语义:商品冻结契约前置(含商品文字锚定),防止风格词改商品
|
||||
if provider_name == "rightapi":
|
||||
prompt = wrap_prompt_for_gpt_edits(prompt)
|
||||
if product:
|
||||
refs = await _select_ref_images(db, product.id, type_id)
|
||||
else:
|
||||
refs = _refs_for_job(list(suite.ref_images or []), job)
|
||||
prompt = wrap_prompt_for_gpt_edits(prompt, ctx)
|
||||
refs = _refs_for_job(list(task.ref_images or []), job)
|
||||
data = await generator(prompt, refs, size=size, model=model)
|
||||
# 部分中转不遵守 output_format(要 jpeg 回 PNG),按魔数定扩展名
|
||||
ext = ".png" if data[:8] == b"\x89PNG\r\n\x1a\n" else ".jpg"
|
||||
key = storage.write_bytes(data, key_prefix=f"suites/{suite.id}", ext=ext)
|
||||
image_row.stored_url = storage.public_url(key)
|
||||
image_row.status = STATUS_OK
|
||||
key = storage.write_bytes(data, key_prefix=f"suites/{task.id}", ext=ext)
|
||||
image.url = storage.public_url(key)
|
||||
image.status = IMG_OK
|
||||
ok += 1
|
||||
except Exception as exc: # noqa: BLE001
|
||||
log.exception("套图 %s 类型 %s 生成失败", suite_id, type_id)
|
||||
err = str(exc)[:500]
|
||||
image_row.error = err
|
||||
failures.append(f"{job.get('title') or type_name(type_id)}:{err[:200]}")
|
||||
log.exception("套图 %s 类型 %s 生成失败", task.id, type_id)
|
||||
image.error = str(exc)[:500]
|
||||
failures.append(f"{job.get('title') or type_name(type_id)}:{str(exc)[:200]}")
|
||||
failed += 1
|
||||
await db.commit()
|
||||
|
||||
suite.status = SUITE_DONE if failed == 0 else (SUITE_PARTIAL if ok > 0 else SUITE_FAILED)
|
||||
task.status = TASK_DONE if failed == 0 else (TASK_PARTIAL if ok > 0 else TASK_FAILED)
|
||||
if failed:
|
||||
uniq = list(dict.fromkeys(failures)) # 去重保序
|
||||
detail = ";".join(uniq[:6])
|
||||
if len(uniq) > 6:
|
||||
detail += f";…等共 {failed} 张失败"
|
||||
if ok == 0:
|
||||
suite.error = f"全部生成失败。{detail}"
|
||||
task.error = f"全部生成失败。{detail}"
|
||||
else:
|
||||
suite.error = f"部分生成失败({failed} 张)。{detail}"
|
||||
from datetime import datetime, timezone
|
||||
suite.finished_at = datetime.now(timezone.utc)
|
||||
if product:
|
||||
product.stage = "generated" # 商品路径才有的阶段升级
|
||||
await db.commit()
|
||||
task.error = f"部分生成失败({failed} 张)。{detail}"
|
||||
|
||||
Reference in New Issue
Block a user