加了一堆模型和一堆功能
This commit is contained in:
@@ -11,7 +11,7 @@ from openai import AsyncOpenAI
|
||||
|
||||
from app.agent.tools import TOOL_DEFINITIONS, execute_tool
|
||||
from app.config import (
|
||||
get_llm_model,
|
||||
get_llm_model_config,
|
||||
get_llm_max_iterations,
|
||||
get_image_model_config,
|
||||
get_max_recent_turns,
|
||||
@@ -22,9 +22,19 @@ from app.services.image_gen import to_data_uri
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _get_client() -> AsyncOpenAI:
|
||||
base_url = os.getenv("OPENAI_BASE_URL")
|
||||
return AsyncOpenAI(base_url=base_url) if base_url else AsyncOpenAI()
|
||||
def _get_client(provider: str) -> AsyncOpenAI:
|
||||
"""按 provider 创建对应的 OpenAI 兼容客户端。"""
|
||||
if provider == "vectorengine":
|
||||
return AsyncOpenAI(
|
||||
api_key=os.getenv("VECTORENGINE_API_KEY"),
|
||||
base_url=os.getenv("VECTORENGINE_BASE_URL", "https://api.vectorengine.ai/v1"),
|
||||
)
|
||||
if provider == "deepseek":
|
||||
return AsyncOpenAI(
|
||||
api_key=os.getenv("DEEPSEEK_API_KEY"),
|
||||
base_url=os.getenv("DEEPSEEK_BASE_URL", "https://api.deepseek.com"),
|
||||
)
|
||||
return AsyncOpenAI()
|
||||
|
||||
SYSTEM_PROMPT = """\
|
||||
你是一个专业的游戏美术 AI 助手。你的工作是帮助美术人员通过对话生成游戏美术资源。
|
||||
@@ -33,7 +43,8 @@ SYSTEM_PROMPT = """\
|
||||
- 根据用户的文字描述生成图片(UI图标、按钮、插画、立绘、概念图等)
|
||||
- 理解用户的审美意图,将中文描述转化为高质量的英文生成 prompt
|
||||
- 根据用户反馈迭代修改(调整颜色、风格、构图等)
|
||||
- 如果用户提供了参考图,将参考图的风格元素融入生成 prompt
|
||||
- 如果用户提供了参考图(支持多张),将参考图的风格元素融入生成 prompt
|
||||
- 理解多张参考图各自的角色(如"图1的主体 + 图2的风格/视角"),并在 prompt 中准确传达
|
||||
- 理解用户在图片上的标注(框选区域 + 文字批注),精准定位需要修改的部分
|
||||
|
||||
## 工作流程
|
||||
@@ -54,23 +65,25 @@ SYSTEM_PROMPT = """\
|
||||
- 必须使用英文
|
||||
- 尽量详细描述:主体内容、颜色方案、光照、构图、材质等
|
||||
- 如果用户要求游戏 UI 元素,添加相关关键词如 "game UI", "icon", "button" 等
|
||||
- 如果你能直接看到参考图(图片内容),可以在 prompt 中描述参考图的风格特征
|
||||
- 如果你无法看到参考图(只收到了文字提示说有参考图),参考图会由生图工具的 IP-Adapter 自动处理风格融合。此时你不要自行猜测画风/艺术风格关键词(如 pixel art、watercolor、oil painting 等),把风格交给参考图来决定。但如果用户在消息中明确指定了风格(如"赛博朋克风"、"水彩风"等),应保留并翻译到 prompt 中——尊重用户的主动意图
|
||||
- 如果你能直接看到参考图(图片内容),可以在 prompt 中描述参考图的风格特征。多张参考图时,理解用户对各图的定位(如"图1做主体参考、图2做风格参考"),将相应特征分别融入 prompt
|
||||
- 如果你无法看到参考图(只收到了文字提示说有参考图),参考图会由生图工具自动处理风格融合。此时你不要自行猜测画风/艺术风格关键词(如 pixel art、watercolor、oil painting 等),把风格交给参考图来决定。但如果用户在消息中明确指定了风格(如"赛博朋克风"、"水彩风"等),应保留并翻译到 prompt 中——尊重用户的主动意图
|
||||
|
||||
## 注意事项
|
||||
- 用中文和用户交流
|
||||
- 生成图片后简要说明你使用的 prompt 思路
|
||||
- 主动建议迭代方向
|
||||
- **禁止模拟工具调用**:生成图片时必须实际调用 generate_image 工具,绝不能用文字描述"已生成"或假装工具已执行。如果需要生成多张图片,每张都必须单独调用工具
|
||||
- **禁止在回复中嵌入图片链接**:不要在回复文字中使用 Markdown 图片语法(如 `` 或 `sandbox:` 链接)。图片展示由系统自动处理,你只需用文字描述结果即可
|
||||
"""
|
||||
|
||||
|
||||
async def run_agent_loop(
|
||||
messages: list[dict],
|
||||
ref_image_url: Optional[str] = None,
|
||||
ref_image_urls: Optional[list[str]] = None,
|
||||
image_model: Optional[str] = None,
|
||||
session_id: Optional[str] = None,
|
||||
user_id: str = "default_user",
|
||||
llm_model: Optional[str] = None,
|
||||
) -> AsyncGenerator[dict, None]:
|
||||
"""
|
||||
运行 Agent Loop,以 SSE 事件流形式 yield 结果。
|
||||
@@ -114,10 +127,16 @@ async def run_agent_loop(
|
||||
yield {"type": "error", "data": {"message": f"记忆系统检索失败: {e}"}}
|
||||
return
|
||||
|
||||
# ── 解析 LLM 模型配置 ──
|
||||
llm_config = get_llm_model_config(llm_model)
|
||||
llm_provider = llm_config["provider"]
|
||||
llm_model_id = llm_config["model_id"]
|
||||
vision_capable = llm_config.get("vision", False)
|
||||
|
||||
# ── 构建 system prompt ──
|
||||
model_config = get_image_model_config(image_model)
|
||||
current_model_name = model_config.get('name', '未知')
|
||||
current_model_id = model_config.get('id', '未知')
|
||||
img_model_config = get_image_model_config(image_model)
|
||||
current_model_name = img_model_config.get('name', '未知')
|
||||
current_model_id = img_model_config.get('id', '未知')
|
||||
model_hint = (
|
||||
f"\n\n## 当前生图模型(重要)\n"
|
||||
f"本次对话用户选择的生图模型是 **{current_model_name}**"
|
||||
@@ -127,43 +146,39 @@ async def run_agent_loop(
|
||||
)
|
||||
api_messages = [{"role": "system", "content": SYSTEM_PROMPT + model_hint + memory_block}]
|
||||
|
||||
# 检测当前 LLM 是否支持 vision(多模态图片输入)
|
||||
llm_model = get_llm_model().lower()
|
||||
vision_capable = any(kw in llm_model for kw in ("gpt-4o", "gpt-4-vision", "claude"))
|
||||
effective_refs = ref_image_urls or []
|
||||
|
||||
for msg in messages:
|
||||
if msg["role"] == "user" and ref_image_url and msg is messages[-1]:
|
||||
if msg["role"] == "user" and effective_refs and msg is messages[-1]:
|
||||
if vision_capable:
|
||||
api_messages.append({
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": msg["content"]},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": to_data_uri(ref_image_url)},
|
||||
},
|
||||
],
|
||||
})
|
||||
content_parts: list[dict] = [{"type": "text", "text": msg["content"]}]
|
||||
for ref_url in effective_refs:
|
||||
content_parts.append({
|
||||
"type": "image_url",
|
||||
"image_url": {"url": to_data_uri(ref_url)},
|
||||
})
|
||||
api_messages.append({"role": "user", "content": content_parts})
|
||||
else:
|
||||
n_refs = len(effective_refs)
|
||||
hint = (
|
||||
f"{msg['content']}\n\n"
|
||||
"【系统提示:用户上传了一张参考图,已自动传递给图片生成工具的 IP-Adapter。"
|
||||
"IP-Adapter 会从参考图中提取风格并融合到生成结果中。"
|
||||
"你无法看到这张参考图,因此在生成 prompt 时:\n"
|
||||
f"【系统提示:用户上传了 {n_refs} 张参考图,已自动传递给图片生成工具。"
|
||||
"生图工具会根据模型能力自动处理参考图的风格融合。"
|
||||
"你无法看到这些参考图,因此在生成 prompt 时:\n"
|
||||
"1. 描述画面内容(主体、构图、光照、材质等)\n"
|
||||
"2. 不要自行猜测画风/艺术风格——但如果用户明确指定了风格,保留到 prompt 中\n"
|
||||
"3. 用户未指定风格时,风格完全由参考图通过 IP-Adapter 决定】"
|
||||
"3. 用户未指定风格时,风格完全由参考图决定】"
|
||||
)
|
||||
api_messages.append({"role": "user", "content": hint})
|
||||
continue
|
||||
api_messages.append({"role": msg["role"], "content": msg["content"]})
|
||||
|
||||
client = _get_client()
|
||||
client = _get_client(llm_provider)
|
||||
|
||||
for _ in range(get_llm_max_iterations()):
|
||||
try:
|
||||
response = await client.chat.completions.create(
|
||||
model=get_llm_model(),
|
||||
model=llm_model_id,
|
||||
messages=api_messages,
|
||||
tools=TOOL_DEFINITIONS,
|
||||
stream=True,
|
||||
@@ -219,7 +234,7 @@ async def run_agent_loop(
|
||||
"你刚才没有调用 generate_image 工具,只是用文字描述了生成过程。"
|
||||
"请立即调用 generate_image 工具来实际生成图片。"
|
||||
"不要解释,直接调用工具。"
|
||||
f"当前使用的生图模型是 {model_config.get('name', '未知')}。"
|
||||
f"当前使用的生图模型是 {img_model_config.get('name', '未知')}。"
|
||||
),
|
||||
})
|
||||
continue
|
||||
@@ -263,7 +278,7 @@ async def run_agent_loop(
|
||||
except json.JSONDecodeError:
|
||||
arguments = {}
|
||||
|
||||
result = await execute_tool(tool_name, arguments, ref_image_url, image_model)
|
||||
result = await execute_tool(tool_name, arguments, effective_refs or None, image_model)
|
||||
|
||||
used_model = result.get("model_name", "")
|
||||
|
||||
|
||||
@@ -11,6 +11,8 @@ TOOL_DEFINITIONS = [
|
||||
"description": (
|
||||
"根据文字描述生成图片。prompt 必须是英文。"
|
||||
"如果用户提供了参考图,会自动传入 ref_image_url 参数。"
|
||||
"若当前生图模型为 Stable Diffusion XL,可在正提示后单独一行写 ---NEGATIVE--- 再写负向提示;"
|
||||
"不传则服务端会使用该模型的默认负向词。"
|
||||
),
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
@@ -37,7 +39,7 @@ TOOL_DEFINITIONS = [
|
||||
async def execute_tool(
|
||||
tool_name: str,
|
||||
arguments: dict,
|
||||
ref_image_url: str | None = None,
|
||||
ref_image_urls: list[str] | None = None,
|
||||
image_model: str | None = None,
|
||||
) -> dict:
|
||||
"""执行工具调用,返回结果。"""
|
||||
@@ -47,7 +49,7 @@ async def execute_tool(
|
||||
result = await generate_images(
|
||||
prompt=prompt,
|
||||
num_images=num_images,
|
||||
ref_image_url=ref_image_url,
|
||||
ref_image_urls=ref_image_urls,
|
||||
model_id=image_model,
|
||||
)
|
||||
valid_urls = [u for u in result.urls if not u.startswith("[")]
|
||||
@@ -57,6 +59,8 @@ async def execute_tool(
|
||||
"images": valid_urls,
|
||||
"errors": errors,
|
||||
"prompt_used": prompt,
|
||||
"effective_prompt": result.effective_prompt,
|
||||
"negative_prompt": result.negative_prompt,
|
||||
"model_name": result.model_name,
|
||||
"model_id": result.model_id,
|
||||
}
|
||||
|
||||
@@ -8,7 +8,12 @@ from sse_starlette.sse import EventSourceResponse
|
||||
|
||||
from app.agent.loop import run_agent_loop
|
||||
from app.auth import get_current_user
|
||||
from app.config import get_image_models_list, get_default_image_model_id
|
||||
from app.config import (
|
||||
get_image_models_list,
|
||||
get_default_image_model_id,
|
||||
get_llm_models_list,
|
||||
get_default_llm_model_id,
|
||||
)
|
||||
from app.db import User
|
||||
|
||||
router = APIRouter()
|
||||
@@ -49,13 +54,24 @@ async def list_models(current_user: User = Depends(get_current_user)):
|
||||
}
|
||||
|
||||
|
||||
@router.get("/llm-models")
|
||||
async def list_llm_models(current_user: User = Depends(get_current_user)):
|
||||
"""返回可用的 LLM 对话模型列表。"""
|
||||
return {
|
||||
"models": get_llm_models_list(),
|
||||
"default": get_default_llm_model_id(),
|
||||
}
|
||||
|
||||
|
||||
@router.post("/chat")
|
||||
async def chat(
|
||||
messages: str = Form(...),
|
||||
ref_image: Optional[UploadFile] = File(None),
|
||||
ref_image_url: Optional[str] = Form(None),
|
||||
ref_image_urls: Optional[str] = Form(None),
|
||||
image_model: Optional[str] = Form(None),
|
||||
session_id: Optional[str] = Form(None),
|
||||
llm_model: Optional[str] = Form(None),
|
||||
current_user: User = Depends(get_current_user),
|
||||
):
|
||||
"""
|
||||
@@ -64,25 +80,30 @@ async def chat(
|
||||
参数:
|
||||
- messages: JSON 字符串,对话历史 [{role, content}]
|
||||
- ref_image: 可选的参考图文件(兼容旧方式)
|
||||
- ref_image_url: 可选,已通过 /upload-ref-image 上传后的服务端路径
|
||||
- ref_image_url: 兼容旧方式,单张参考图服务端路径
|
||||
- ref_image_urls: JSON 数组字符串,多张参考图服务端路径列表
|
||||
- image_model: 可选,指定本次使用的生图模型短 ID
|
||||
- session_id: 可选,前端会话 ID,用于 Mem0 记忆作用域
|
||||
- llm_model: 可选,指定本次使用的 LLM 模型短 ID
|
||||
"""
|
||||
parsed_messages = json.loads(messages)
|
||||
|
||||
resolved_ref_url: Optional[str] = None
|
||||
if ref_image_url:
|
||||
resolved_ref_url = ref_image_url
|
||||
resolved_ref_urls: list[str] = []
|
||||
if ref_image_urls:
|
||||
resolved_ref_urls = json.loads(ref_image_urls)
|
||||
elif ref_image_url:
|
||||
resolved_ref_urls = [ref_image_url]
|
||||
elif ref_image and ref_image.filename:
|
||||
resolved_ref_url = await _save_upload(ref_image)
|
||||
resolved_ref_urls = [await _save_upload(ref_image)]
|
||||
|
||||
async def event_generator():
|
||||
async for event in run_agent_loop(
|
||||
parsed_messages,
|
||||
resolved_ref_url,
|
||||
resolved_ref_urls or None,
|
||||
image_model=image_model,
|
||||
session_id=session_id,
|
||||
user_id=current_user.id,
|
||||
llm_model=llm_model,
|
||||
):
|
||||
yield {
|
||||
"event": event["type"],
|
||||
|
||||
37
art-agent/backend/app/api/memory.py
Normal file
37
art-agent/backend/app/api/memory.py
Normal file
@@ -0,0 +1,37 @@
|
||||
"""
|
||||
记忆查询 API:供前端查看当前用户的 Mem0 记忆条目。
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
|
||||
from fastapi import APIRouter, Depends
|
||||
|
||||
from app.auth import get_current_user
|
||||
from app.db import User
|
||||
from app.memory import get_memory
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/memory", tags=["memory"])
|
||||
|
||||
|
||||
@router.get("/list")
|
||||
async def list_memories(current_user: User = Depends(get_current_user)):
|
||||
"""返回当前用户的所有记忆条目(只读)。"""
|
||||
memory = get_memory()
|
||||
loop = asyncio.get_running_loop()
|
||||
try:
|
||||
result = await loop.run_in_executor(
|
||||
None, lambda: memory.get_all(user_id=current_user.id, limit=200)
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error("获取记忆列表失败: %s", e)
|
||||
return {"memories": [], "error": str(e)}
|
||||
|
||||
results = result.get("results", []) if isinstance(result, dict) else result
|
||||
logger.info("Mem0 get_all 返回 %d 条记忆, result_type=%s", len(results), type(result).__name__)
|
||||
if results:
|
||||
sample = results[0]
|
||||
logger.info("记忆样本 keys=%s, memory=%s", list(sample.keys()) if isinstance(sample, dict) else "not-dict", str(sample)[:200])
|
||||
return {"memories": results}
|
||||
@@ -7,14 +7,108 @@ import os
|
||||
from typing import Any
|
||||
|
||||
|
||||
def get_llm_model() -> str:
|
||||
return os.getenv("LLM_MODEL", "gpt-4o-mini")
|
||||
|
||||
|
||||
def get_llm_max_iterations() -> int:
|
||||
return int(os.getenv("LLM_MAX_ITERATIONS", "5"))
|
||||
|
||||
|
||||
# ─── LLM 模型注册表 ─────────────────────────────────────
|
||||
#
|
||||
# 每个模型的配置说明:
|
||||
# id — 前端/API 使用的短 ID
|
||||
# name — 显示名称
|
||||
# provider — API 提供者(对应 _get_client 的分发键)
|
||||
# model_id — 传给 OpenAI SDK 的 model 参数
|
||||
# description — 前端下拉列表说明文字
|
||||
# vision — 是否支持多模态图片输入
|
||||
|
||||
LLM_MODELS: dict[str, dict[str, Any]] = {
|
||||
"gpt-5.4": {
|
||||
"id": "gpt-5.4",
|
||||
"name": "GPT-5.4",
|
||||
"provider": "vectorengine",
|
||||
"model_id": "gpt-5.4",
|
||||
"description": "知识工作与计算机操控最强,1M 上下文",
|
||||
"vision": True,
|
||||
},
|
||||
"claude-sonnet-4-6": {
|
||||
"id": "claude-sonnet-4-6",
|
||||
"name": "Claude Sonnet 4.6",
|
||||
"provider": "vectorengine",
|
||||
"model_id": "claude-sonnet-4-6",
|
||||
"description": "高性价比编码与日常任务",
|
||||
"vision": True,
|
||||
},
|
||||
"claude-opus-4-6": {
|
||||
"id": "claude-opus-4-6",
|
||||
"name": "Claude Opus 4.6",
|
||||
"provider": "vectorengine",
|
||||
"model_id": "claude-opus-4-6",
|
||||
"description": "编码与专家级推理最强,128K 输出",
|
||||
"vision": True,
|
||||
},
|
||||
"gemini-3.1-pro-preview": {
|
||||
"id": "gemini-3.1-pro-preview",
|
||||
"name": "Gemini 3.1 Pro",
|
||||
"provider": "vectorengine",
|
||||
"model_id": "gemini-3.1-pro-preview",
|
||||
"description": "推理最强、价格最低,2M 上下文",
|
||||
"vision": True,
|
||||
},
|
||||
"glm-4.7": {
|
||||
"id": "glm-4.7",
|
||||
"name": "GLM-4.7",
|
||||
"provider": "vectorengine",
|
||||
"model_id": "glm-4.7",
|
||||
"description": "智谱 AI,中文能力突出,免费额度",
|
||||
"vision": False,
|
||||
},
|
||||
"gpt-4o-mini": {
|
||||
"id": "gpt-4o-mini",
|
||||
"name": "GPT-4o Mini",
|
||||
"provider": "vectorengine",
|
||||
"model_id": "gpt-4o-mini",
|
||||
"description": "轻量快速、高性价比",
|
||||
"vision": True,
|
||||
},
|
||||
"deepseek-chat": {
|
||||
"id": "deepseek-chat",
|
||||
"name": "DeepSeek Chat",
|
||||
"provider": "deepseek",
|
||||
"model_id": "deepseek-chat",
|
||||
"description": "中文对话优化(直连)",
|
||||
"vision": False,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def get_default_llm_model_id() -> str:
|
||||
"""返回 .env 中配置的默认 LLM 短 ID,不在注册表中则回退到 gpt-5.4。"""
|
||||
env_model = os.getenv("LLM_MODEL", "gpt-5.4")
|
||||
if env_model in LLM_MODELS:
|
||||
return env_model
|
||||
return "gpt-5.4"
|
||||
|
||||
|
||||
def get_llm_model_config(model_id: str | None = None) -> dict[str, Any]:
|
||||
"""根据短 ID 获取 LLM 模型配置,未指定或不存在则使用默认模型。"""
|
||||
if model_id and model_id in LLM_MODELS:
|
||||
return LLM_MODELS[model_id]
|
||||
return LLM_MODELS[get_default_llm_model_id()]
|
||||
|
||||
|
||||
def get_llm_models_list() -> list[dict]:
|
||||
"""返回前端下拉列表所需的 LLM 模型摘要信息。"""
|
||||
return [
|
||||
{
|
||||
"id": cfg["id"],
|
||||
"name": cfg["name"],
|
||||
"description": cfg["description"],
|
||||
"vision": cfg.get("vision", False),
|
||||
}
|
||||
for cfg in LLM_MODELS.values()
|
||||
]
|
||||
|
||||
|
||||
# ─── 图像模型注册表 ─────────────────────────────────────
|
||||
#
|
||||
# 每个模型的配置说明:
|
||||
@@ -29,6 +123,29 @@ def get_llm_max_iterations() -> int:
|
||||
# default_params — 默认推理参数
|
||||
|
||||
IMAGE_MODELS: dict[str, dict[str, Any]] = {
|
||||
"gpt-image-1.5": {
|
||||
"id": "gpt-image-1.5",
|
||||
"name": "GPT Image 1.5",
|
||||
"provider": "openai",
|
||||
"model_id": "gpt-image-1.5",
|
||||
"description": "OpenAI 最强生图,文字渲染与 prompt 理解最佳,支持多图参考",
|
||||
"supports_ref_image": True,
|
||||
"max_ref_images": 16,
|
||||
"default_params": {
|
||||
"size": "1024x1024",
|
||||
"quality": "high",
|
||||
},
|
||||
},
|
||||
"gemini-3.1-flash-image": {
|
||||
"id": "gemini-3.1-flash-image",
|
||||
"name": "Gemini 3.1 Flash Image",
|
||||
"provider": "gemini_native",
|
||||
"model_id": "gemini-3.1-flash-image-preview",
|
||||
"description": "Google 原生生图,速度快、价格低,支持多图参考(最多 14 张)",
|
||||
"supports_ref_image": True,
|
||||
"max_ref_images": 14,
|
||||
"default_params": {},
|
||||
},
|
||||
"flux-schnell": {
|
||||
"id": "flux-schnell",
|
||||
"name": "Flux Schnell",
|
||||
|
||||
@@ -12,6 +12,7 @@ from fastapi.staticfiles import StaticFiles
|
||||
from app.api.chat import router as chat_router
|
||||
from app.api.auth import router as auth_router
|
||||
from app.api.admin import router as admin_router
|
||||
from app.api.memory import router as memory_router
|
||||
from app.db import create_db_and_tables, ensure_default_admin
|
||||
|
||||
|
||||
@@ -45,6 +46,7 @@ app.mount("/generated", StaticFiles(directory=str(GENERATED_DIR)), name="generat
|
||||
app.include_router(auth_router, prefix="/api")
|
||||
app.include_router(admin_router, prefix="/api")
|
||||
app.include_router(chat_router, prefix="/api")
|
||||
app.include_router(memory_router, prefix="/api")
|
||||
|
||||
|
||||
@app.get("/health")
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
"""图像生成服务 — Provider 抽象层。"""
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import mimetypes
|
||||
import os
|
||||
@@ -9,9 +10,11 @@ from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from openai import AsyncOpenAI
|
||||
from replicate import Client as ReplicateClient
|
||||
|
||||
from app.config import get_image_model_config, get_image_aspect_ratio, get_image_output_format
|
||||
from app.config import get_image_aspect_ratio, get_image_model_config, get_image_output_format
|
||||
from app.services.image_prompt_strategy import prepare_image_prompt_for_model
|
||||
|
||||
BACKEND_ROOT = Path(__file__).parent.parent.parent
|
||||
GENERATED_DIR = BACKEND_ROOT / "generated"
|
||||
@@ -57,6 +60,43 @@ def to_data_uri(image_path: str) -> str:
|
||||
return image_path
|
||||
|
||||
|
||||
def _load_image_bytes(image_path: str) -> bytes:
|
||||
"""将本地路径或 data URI 转为原始字节,供 OpenAI images.edit 使用。"""
|
||||
if image_path.startswith("data:"):
|
||||
# data:image/png;base64,xxxx
|
||||
_, b64_part = image_path.split(",", 1)
|
||||
return base64.b64decode(b64_part)
|
||||
if image_path.startswith("http"):
|
||||
raise ValueError("_load_image_bytes 不支持远程 URL,请先下载到本地")
|
||||
local = BACKEND_ROOT / image_path.lstrip("/")
|
||||
if local.exists():
|
||||
return local.read_bytes()
|
||||
raise FileNotFoundError(f"参考图文件未找到: {local}")
|
||||
|
||||
|
||||
def _resolve_image_base64(image_path: str) -> tuple[str, str]:
|
||||
"""将图片路径/data URI 解析为 (mime_type, base64_string)。
|
||||
|
||||
统一处理三种输入形式:
|
||||
- data URI (data:image/png;base64,xxxx) → 直接提取 mime 和 base64
|
||||
- 本地路径 (/uploads/xxx.png) → 读取文件并编码
|
||||
- 其他 → 尝试作为本地路径处理
|
||||
"""
|
||||
if image_path.startswith("data:"):
|
||||
header, b64_part = image_path.split(",", 1)
|
||||
# header 格式: data:image/png;base64
|
||||
mime = header.split(";")[0].replace("data:", "")
|
||||
return mime, b64_part
|
||||
|
||||
local = BACKEND_ROOT / image_path.lstrip("/")
|
||||
if local.exists():
|
||||
mime = mimetypes.guess_type(str(local))[0] or "image/png"
|
||||
b64 = base64.b64encode(local.read_bytes()).decode()
|
||||
return mime, b64
|
||||
|
||||
raise FileNotFoundError(f"参考图文件未找到: {local}")
|
||||
|
||||
|
||||
async def _download_image(url: str) -> str:
|
||||
"""下载远程图片到本地 generated/ 目录,返回本地 URL 路径。"""
|
||||
filename = f"{uuid.uuid4().hex}.png"
|
||||
@@ -80,9 +120,10 @@ class ImageProvider(ABC):
|
||||
model_config: dict[str, Any],
|
||||
prompt: str,
|
||||
num_images: int = 1,
|
||||
ref_image_url: str | None = None,
|
||||
ref_image_urls: list[str] | None = None,
|
||||
negative_prompt: str | None = None,
|
||||
) -> list[str]:
|
||||
"""生成图片并返回本地 URL 列表。"""
|
||||
"""生成图片并返回本地 URL 列表。negative_prompt 仅部分后端使用(如 Replicate SDXL)。"""
|
||||
...
|
||||
|
||||
|
||||
@@ -94,7 +135,8 @@ class ReplicateProvider(ImageProvider):
|
||||
model_config: dict[str, Any],
|
||||
prompt: str,
|
||||
num_images: int = 1,
|
||||
ref_image_url: str | None = None,
|
||||
ref_image_urls: list[str] | None = None,
|
||||
negative_prompt: str | None = None,
|
||||
) -> list[str]:
|
||||
replicate_model_id = model_config["model_id"]
|
||||
default_params: dict[str, Any] = model_config.get("default_params", {})
|
||||
@@ -102,8 +144,15 @@ class ReplicateProvider(ImageProvider):
|
||||
local_urls: list[str] = []
|
||||
|
||||
try:
|
||||
# Replicate 模型(InstantStyle/Kolors)只支持单张参考图,取第一张
|
||||
single_ref = ref_image_urls[0] if ref_image_urls else None
|
||||
input_params = self._build_input(
|
||||
model_config, default_params, prompt, num_images, ref_image_url
|
||||
model_config,
|
||||
default_params,
|
||||
prompt,
|
||||
num_images,
|
||||
single_ref,
|
||||
negative_prompt=negative_prompt,
|
||||
)
|
||||
|
||||
output = await _replicate_client.async_run(
|
||||
@@ -131,20 +180,18 @@ class ReplicateProvider(ImageProvider):
|
||||
prompt: str,
|
||||
num_images: int = 1,
|
||||
ref_image_url: str | None = None,
|
||||
negative_prompt: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
根据模型配置构建 Replicate input 参数。
|
||||
|
||||
通过 model_config 中的标志字段自动适配不同模型:
|
||||
- supports_ref_image / ref_image_param: 参考图注入
|
||||
- num_images_param: 各模型的批量生成参数名(如 number_of_images / num_outputs)
|
||||
Replicate 上的 IP-Adapter 模型只支持单张参考图。
|
||||
"""
|
||||
model_id = model_config["model_id"]
|
||||
supports_ref = model_config.get("supports_ref_image", False)
|
||||
ref_param_name = model_config.get("ref_image_param", "image")
|
||||
num_images_param = model_config.get("num_images_param", "num_outputs")
|
||||
|
||||
# ── 支持参考图的模型(IP-Adapter 系列)──
|
||||
if supports_ref:
|
||||
params: dict[str, Any] = {"prompt": prompt}
|
||||
for k, v in default_params.items():
|
||||
@@ -154,7 +201,6 @@ class ReplicateProvider(ImageProvider):
|
||||
params[ref_param_name] = to_data_uri(ref_image_url)
|
||||
return params
|
||||
|
||||
# ── Flux 系列(纯文生图)──
|
||||
is_flux = "flux" in model_id.lower()
|
||||
params = {"prompt": prompt, num_images_param: num_images}
|
||||
|
||||
@@ -166,21 +212,281 @@ class ReplicateProvider(ImageProvider):
|
||||
"output_format", get_image_output_format()
|
||||
)
|
||||
else:
|
||||
# SDXL 类模型
|
||||
params["width"] = default_params.get("width", 1024)
|
||||
params["height"] = default_params.get("height", 1024)
|
||||
if "num_inference_steps" in default_params:
|
||||
params["num_inference_steps"] = default_params["num_inference_steps"]
|
||||
if "guidance_scale" in default_params:
|
||||
params["guidance_scale"] = default_params["guidance_scale"]
|
||||
if negative_prompt is not None:
|
||||
params["negative_prompt"] = negative_prompt
|
||||
|
||||
return params
|
||||
|
||||
|
||||
class OpenAIImageProvider(ImageProvider):
|
||||
"""通过 OpenAI 兼容 API(向量引擎中转)调用 GPT Image 系列模型。
|
||||
|
||||
有参考图时使用 images.edit(支持最多 16 张参考图),
|
||||
无参考图时使用 images.generate(纯文生图)。
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._client: AsyncOpenAI | None = None
|
||||
|
||||
def _get_client(self) -> AsyncOpenAI:
|
||||
if self._client is None:
|
||||
self._client = AsyncOpenAI(
|
||||
api_key=os.getenv("VECTORENGINE_API_KEY"),
|
||||
base_url=os.getenv("VECTORENGINE_BASE_URL", "https://api.vectorengine.ai/v1"),
|
||||
timeout=httpx.Timeout(5.0, read=180.0, write=60.0, connect=30.0),
|
||||
)
|
||||
return self._client
|
||||
|
||||
async def generate(
|
||||
self,
|
||||
model_config: dict[str, Any],
|
||||
prompt: str,
|
||||
num_images: int = 1,
|
||||
ref_image_urls: list[str] | None = None,
|
||||
negative_prompt: str | None = None,
|
||||
) -> list[str]:
|
||||
_ = negative_prompt
|
||||
client = self._get_client()
|
||||
default_params = model_config.get("default_params", {})
|
||||
model_id = model_config["model_id"]
|
||||
|
||||
local_urls: list[str] = []
|
||||
try:
|
||||
if ref_image_urls:
|
||||
resp = await self._edit_with_refs(
|
||||
client, model_id, prompt, ref_image_urls, num_images, default_params
|
||||
)
|
||||
else:
|
||||
resp = await client.images.generate(
|
||||
model=model_id,
|
||||
prompt=prompt,
|
||||
n=num_images,
|
||||
size=default_params.get("size", "1024x1024"),
|
||||
quality=default_params.get("quality", "high"),
|
||||
)
|
||||
|
||||
for img_data in resp.data:
|
||||
if img_data.url:
|
||||
local_urls.append(await _download_image(img_data.url))
|
||||
elif img_data.b64_json:
|
||||
filename = f"{uuid.uuid4().hex}.png"
|
||||
filepath = GENERATED_DIR / filename
|
||||
filepath.write_bytes(base64.b64decode(img_data.b64_json))
|
||||
local_urls.append(f"/generated/{filename}")
|
||||
else:
|
||||
local_urls.append("[生成失败: 模型未返回图片数据]")
|
||||
|
||||
except Exception as e:
|
||||
detail = str(e) or f"{type(e).__name__}: {repr(e)}"
|
||||
local_urls.append(f"[生成失败: {detail}]")
|
||||
|
||||
return local_urls
|
||||
|
||||
@staticmethod
|
||||
async def _edit_with_refs(
|
||||
client: AsyncOpenAI,
|
||||
model_id: str,
|
||||
prompt: str,
|
||||
ref_image_urls: list[str],
|
||||
num_images: int,
|
||||
default_params: dict[str, Any],
|
||||
):
|
||||
"""使用 images.edit 端点传入参考图(GPT Image 系列最多 16 张)。"""
|
||||
image_files: list[Any] = []
|
||||
for url in ref_image_urls:
|
||||
image_files.append(_load_image_bytes(url))
|
||||
|
||||
image_arg: Any = image_files[0] if len(image_files) == 1 else image_files
|
||||
|
||||
return await client.images.edit(
|
||||
model=model_id,
|
||||
image=image_arg,
|
||||
prompt=prompt,
|
||||
n=num_images,
|
||||
size=default_params.get("size", "1024x1024"),
|
||||
)
|
||||
|
||||
|
||||
class GeminiNativeImageProvider(ImageProvider):
|
||||
"""通过向量引擎中转调用 Gemini 原生 generateContent 接口。
|
||||
|
||||
Gemini 原生接口支持文字 + 图片混合输入(最多 14 张参考图),
|
||||
在一次 generateContent 调用中同时理解参考图并生成新图片。
|
||||
这是 OpenAI 兼容的 images/generate 和 images/edit 都无法覆盖的能力。
|
||||
|
||||
API 格式:
|
||||
POST /v1beta/models/{model}:generateContent?key={API_KEY}
|
||||
Body: { contents: [{ parts: [...] }], generationConfig: { responseModalities: ["TEXT","IMAGE"] } }
|
||||
Response: candidates[0].content.parts[] → text 或 inline_data (base64)
|
||||
|
||||
多图 + 生图耗时较长,上游或代理可能提前断开(httpx: Server disconnected without sending a response)。
|
||||
使用较长超时 + 对可恢复网络错误自动重试。
|
||||
"""
|
||||
|
||||
_RETRYABLE: tuple[type[BaseException], ...] = (
|
||||
httpx.RemoteProtocolError,
|
||||
httpx.ConnectError,
|
||||
httpx.ReadTimeout,
|
||||
httpx.WriteTimeout,
|
||||
httpx.ConnectTimeout,
|
||||
httpx.PoolTimeout,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _gemini_timeout(cls) -> httpx.Timeout:
|
||||
"""可通过环境变量调大,多参考图时请求体大、响应慢。"""
|
||||
read_s = float(os.getenv("VECTORENGINE_GEMINI_READ_TIMEOUT", "600"))
|
||||
write_s = float(os.getenv("VECTORENGINE_GEMINI_WRITE_TIMEOUT", "180"))
|
||||
connect_s = float(os.getenv("VECTORENGINE_GEMINI_CONNECT_TIMEOUT", "60"))
|
||||
pool_s = float(os.getenv("VECTORENGINE_GEMINI_POOL_TIMEOUT", "60"))
|
||||
return httpx.Timeout(
|
||||
connect=connect_s,
|
||||
read=read_s,
|
||||
write=write_s,
|
||||
pool=pool_s,
|
||||
)
|
||||
|
||||
async def generate(
|
||||
self,
|
||||
model_config: dict[str, Any],
|
||||
prompt: str,
|
||||
num_images: int = 1,
|
||||
ref_image_urls: list[str] | None = None,
|
||||
negative_prompt: str | None = None,
|
||||
) -> list[str]:
|
||||
_ = negative_prompt
|
||||
api_key = os.getenv("VECTORENGINE_API_KEY", "")
|
||||
base_url = os.getenv("VECTORENGINE_BASE_URL", "https://api.vectorengine.ai/v1")
|
||||
# 从 /v1 回退到根 URL,拼接 /v1beta/models/... 端点
|
||||
api_root = base_url.rstrip("/").removesuffix("/v1")
|
||||
model_id = model_config["model_id"]
|
||||
|
||||
url = f"{api_root}/v1beta/models/{model_id}:generateContent?key={api_key}"
|
||||
|
||||
parts = self._build_parts(prompt, ref_image_urls)
|
||||
body = {
|
||||
"contents": [{"parts": parts}],
|
||||
"generationConfig": {
|
||||
"responseModalities": ["TEXT", "IMAGE"],
|
||||
},
|
||||
}
|
||||
|
||||
local_urls: list[str] = []
|
||||
proxy = os.environ.get("HTTPS_PROXY") or os.environ.get("HTTP_PROXY")
|
||||
max_retries = max(1, int(os.getenv("VECTORENGINE_GEMINI_MAX_RETRIES", "3")))
|
||||
timeout = self._gemini_timeout()
|
||||
|
||||
try:
|
||||
data: dict[str, Any] | None = None
|
||||
for attempt in range(max_retries):
|
||||
try:
|
||||
async with httpx.AsyncClient(
|
||||
proxy=proxy,
|
||||
timeout=timeout,
|
||||
limits=httpx.Limits(max_keepalive_connections=5, max_connections=10),
|
||||
) as client:
|
||||
resp = await client.post(
|
||||
url,
|
||||
json=body,
|
||||
headers={"Content-Type": "application/json"},
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
break
|
||||
except self._RETRYABLE as e:
|
||||
if attempt < max_retries - 1:
|
||||
await asyncio.sleep(2 ** attempt)
|
||||
continue
|
||||
raise
|
||||
|
||||
if data is None:
|
||||
local_urls.append("[生成失败: 未收到上游响应]")
|
||||
return local_urls
|
||||
|
||||
local_urls = self._extract_images(data)
|
||||
|
||||
if not local_urls:
|
||||
# Gemini 可能只返回了文字(拒绝生图或纯文字回复)
|
||||
text_parts = self._extract_text(data)
|
||||
hint = text_parts[:200] if text_parts else "模型未返回图片"
|
||||
local_urls.append(f"[生成失败: {hint}]")
|
||||
|
||||
except httpx.HTTPStatusError as e:
|
||||
detail = e.response.text[:500] if e.response else str(e)
|
||||
local_urls.append(f"[生成失败: HTTP {e.response.status_code} - {detail}]")
|
||||
except self._RETRYABLE as e:
|
||||
hint = (
|
||||
"连接被上游或代理提前关闭,常见于多图参考或生图较慢。"
|
||||
"可稍后重试,或在 .env 中增大 VECTORENGINE_GEMINI_READ_TIMEOUT / 检查代理稳定性。"
|
||||
)
|
||||
detail = str(e) or f"{type(e).__name__}: {repr(e)}"
|
||||
local_urls.append(f"[生成失败: {detail}。{hint}]")
|
||||
except Exception as e:
|
||||
detail = str(e) or f"{type(e).__name__}: {repr(e)}"
|
||||
local_urls.append(f"[生成失败: {detail}]")
|
||||
|
||||
return local_urls
|
||||
|
||||
@staticmethod
|
||||
def _build_parts(
|
||||
prompt: str, ref_image_urls: list[str] | None
|
||||
) -> list[dict[str, Any]]:
|
||||
"""构建 Gemini generateContent 的 parts 数组:文字 + 内联图片。"""
|
||||
parts: list[dict[str, Any]] = [{"text": prompt}]
|
||||
if ref_image_urls:
|
||||
for url in ref_image_urls:
|
||||
mime, b64 = _resolve_image_base64(url)
|
||||
parts.append({
|
||||
"inline_data": {
|
||||
"mime_type": mime,
|
||||
"data": b64,
|
||||
}
|
||||
})
|
||||
return parts
|
||||
|
||||
@staticmethod
|
||||
def _extract_images(response_data: dict) -> list[str]:
|
||||
"""从 Gemini 响应中提取所有图片并保存到本地。"""
|
||||
local_urls: list[str] = []
|
||||
candidates = response_data.get("candidates", [])
|
||||
for candidate in candidates:
|
||||
parts = candidate.get("content", {}).get("parts", [])
|
||||
for part in parts:
|
||||
inline = part.get("inlineData") or part.get("inline_data")
|
||||
if inline and inline.get("data"):
|
||||
mime = inline.get("mimeType") or inline.get("mime_type", "image/png")
|
||||
ext = ".png" if "png" in mime else ".jpg" if "jpeg" in mime or "jpg" in mime else ".webp" if "webp" in mime else ".png"
|
||||
filename = f"{uuid.uuid4().hex}{ext}"
|
||||
filepath = GENERATED_DIR / filename
|
||||
filepath.write_bytes(base64.b64decode(inline["data"]))
|
||||
local_urls.append(f"/generated/{filename}")
|
||||
return local_urls
|
||||
|
||||
@staticmethod
|
||||
def _extract_text(response_data: dict) -> str:
|
||||
"""从 Gemini 响应中提取文字内容(用于调试或错误提示)。"""
|
||||
texts: list[str] = []
|
||||
candidates = response_data.get("candidates", [])
|
||||
for candidate in candidates:
|
||||
parts = candidate.get("content", {}).get("parts", [])
|
||||
for part in parts:
|
||||
if part.get("text"):
|
||||
texts.append(part["text"])
|
||||
return "\n".join(texts)
|
||||
|
||||
|
||||
# ─── Provider 注册 ─────────────────────────────────────
|
||||
|
||||
_PROVIDERS: dict[str, ImageProvider] = {
|
||||
"replicate": ReplicateProvider(),
|
||||
"openai": OpenAIImageProvider(),
|
||||
"gemini_native": GeminiNativeImageProvider(),
|
||||
}
|
||||
|
||||
|
||||
@@ -189,21 +495,32 @@ _PROVIDERS: dict[str, ImageProvider] = {
|
||||
class GenerateResult:
|
||||
"""图片生成结果,包含生成的 URL 列表和实际使用的模型信息。"""
|
||||
|
||||
def __init__(self, urls: list[str], model_name: str, model_id: str):
|
||||
def __init__(
|
||||
self,
|
||||
urls: list[str],
|
||||
model_name: str,
|
||||
model_id: str,
|
||||
*,
|
||||
effective_prompt: str | None = None,
|
||||
negative_prompt: str | None = None,
|
||||
):
|
||||
self.urls = urls
|
||||
self.model_name = model_name
|
||||
self.model_id = model_id
|
||||
self.effective_prompt = effective_prompt
|
||||
self.negative_prompt = negative_prompt
|
||||
|
||||
|
||||
async def generate_images(
|
||||
prompt: str,
|
||||
num_images: int = 1,
|
||||
ref_image_url: str | None = None,
|
||||
ref_image_urls: list[str] | None = None,
|
||||
model_id: str | None = None,
|
||||
) -> GenerateResult:
|
||||
"""
|
||||
统一入口:根据 model_id 查注册表,分发到对应 provider。
|
||||
|
||||
ref_image_urls: 参考图路径列表(可为 None 或空列表)。
|
||||
model_id 为空时使用 .env 中配置的默认模型。
|
||||
"""
|
||||
config = get_image_model_config(model_id)
|
||||
@@ -216,5 +533,41 @@ async def generate_images(
|
||||
if not provider:
|
||||
return GenerateResult([f"[未知 provider: {provider_name}]"], model_name, resolved_id)
|
||||
|
||||
urls = await provider.generate(config, prompt, num_images, ref_image_url)
|
||||
return GenerateResult(urls, model_name, resolved_id)
|
||||
effective_refs = ref_image_urls or []
|
||||
|
||||
# Replicate IP-Adapter 系列模型必须有参考图才能工作
|
||||
if config.get("supports_ref_image") and not effective_refs and provider_name == "replicate":
|
||||
return GenerateResult(
|
||||
[f"[生成失败: {model_name} 是风格迁移模型,需要上传参考图才能使用]"],
|
||||
model_name,
|
||||
resolved_id,
|
||||
)
|
||||
|
||||
prepared = prepare_image_prompt_for_model(
|
||||
config,
|
||||
prompt,
|
||||
has_reference_images=bool(effective_refs),
|
||||
)
|
||||
if not prepared.prompt.strip():
|
||||
return GenerateResult(
|
||||
[f"[生成失败: 经模型策略处理后的 prompt 为空]"],
|
||||
model_name,
|
||||
resolved_id,
|
||||
effective_prompt=prepared.prompt,
|
||||
negative_prompt=prepared.negative_prompt,
|
||||
)
|
||||
|
||||
urls = await provider.generate(
|
||||
config,
|
||||
prepared.prompt,
|
||||
num_images,
|
||||
effective_refs or None,
|
||||
negative_prompt=prepared.negative_prompt,
|
||||
)
|
||||
return GenerateResult(
|
||||
urls,
|
||||
model_name,
|
||||
resolved_id,
|
||||
effective_prompt=prepared.prompt,
|
||||
negative_prompt=prepared.negative_prompt,
|
||||
)
|
||||
|
||||
139
art-agent/backend/app/services/image_prompt_strategy.py
Normal file
139
art-agent/backend/app/services/image_prompt_strategy.py
Normal file
@@ -0,0 +1,139 @@
|
||||
"""按生图模型预处理 prompt(与 provider 解耦)。
|
||||
|
||||
各策略在对应函数中注明依据:厂商文档、Replicate 模型 API 字段说明或社区通用写法。
|
||||
统一入口:prepare_image_prompt_for_model。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
# ─── SDXL(Stability / Replicate)────────────────────────────
|
||||
# Replicate stability-ai/sdxl 提供独立字段 negative_prompt(见模型 API 页)。
|
||||
# 下列为 CLIP 系文生图常见「质量/解剖」排除词,作未显式指定时的基线。
|
||||
DEFAULT_SDXL_NEGATIVE = (
|
||||
"low quality, worst quality, normal quality, lowres, blurry, jpeg artifacts, "
|
||||
"watermark, signature, text, logo, deformed, disfigured, bad anatomy, bad hands, "
|
||||
"extra fingers, mutated, cropped, poorly drawn face"
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class PreparedImagePrompt:
|
||||
"""下游只读:prompt 必填;negative_prompt 仅部分后端使用(当前为 SDXL)。"""
|
||||
|
||||
prompt: str
|
||||
negative_prompt: str | None = None
|
||||
|
||||
|
||||
def prepare_image_prompt_for_model(
|
||||
model_config: dict[str, Any],
|
||||
raw_prompt: str,
|
||||
*,
|
||||
has_reference_images: bool = False,
|
||||
) -> PreparedImagePrompt:
|
||||
"""根据注册表 id 选择策略。未知 id 时原样透传。"""
|
||||
short_id = model_config.get("id", "")
|
||||
text = (raw_prompt or "").strip()
|
||||
if not text:
|
||||
return PreparedImagePrompt(prompt="")
|
||||
|
||||
dispatch: dict[str, Any] = {
|
||||
"gpt-image-1.5": _openai_gpt_image,
|
||||
"gemini-3.1-flash-image": _gemini_native_image,
|
||||
"flux-schnell": _flux_bfl,
|
||||
"flux-dev": _flux_bfl,
|
||||
"sdxl": _sdxl_replicate,
|
||||
"instant-style": _replicate_ip_adapter_scene,
|
||||
"kolors-ipadapter": _replicate_ip_adapter_scene,
|
||||
}
|
||||
fn = dispatch.get(short_id, _passthrough)
|
||||
return fn(text, model_config, has_reference_images)
|
||||
|
||||
|
||||
def _passthrough(text: str, _model_config: dict[str, Any], _has_ref: bool) -> PreparedImagePrompt:
|
||||
return PreparedImagePrompt(prompt=text)
|
||||
|
||||
|
||||
def _openai_gpt_image(text: str, _model_config: dict[str, Any], _has_ref: bool) -> PreparedImagePrompt:
|
||||
"""OpenAI GPT Image:自然语言指令遵循强,宜为完整、具体的场景描述。
|
||||
|
||||
参考:https://platform.openai.com/docs/guides/image-generation
|
||||
"""
|
||||
return PreparedImagePrompt(prompt=_collapse_ws(text))
|
||||
|
||||
|
||||
def _gemini_native_image(text: str, _model_config: dict[str, Any], _has_ref: bool) -> PreparedImagePrompt:
|
||||
"""Gemini 原生生图:多模态指令 + 文本;英文描述通常效果稳定。
|
||||
|
||||
参考:Google AI 文档 generateContent 与图像输出 modality 说明。
|
||||
"""
|
||||
return PreparedImagePrompt(prompt=_collapse_ws(text))
|
||||
|
||||
|
||||
def _flux_bfl(text: str, _model_config: dict[str, Any], _has_ref: bool) -> PreparedImagePrompt:
|
||||
"""Black Forest Labs FLUX:Subject + Action + Style + Context;靠前放置重点;无 negative API。
|
||||
|
||||
参考:https://docs.bfl.ai/guides/prompting_guide_t2i_fundamentals
|
||||
若用户从 SDXL 复制了 Negative 段,尽量剥掉以免干扰文意。
|
||||
"""
|
||||
cleaned = _strip_pasted_negative_block(text)
|
||||
return PreparedImagePrompt(prompt=_collapse_ws(cleaned))
|
||||
|
||||
|
||||
def _sdxl_replicate(text: str, model_config: dict[str, Any], _has_ref: bool) -> PreparedImagePrompt:
|
||||
"""Replicate SDXL:prompt + negative_prompt 双字段;支持显式拆分。
|
||||
|
||||
约定(可选):正提示与负提示用单独一行分隔符,便于 Agent/用户手写。
|
||||
- ---NEGATIVE--- 或 |||NEG|||
|
||||
未拆分时使用 default_params.negative_prompt 或模块默认 DEFAULT_SDXL_NEGATIVE。
|
||||
"""
|
||||
neg_fallback = model_config.get("default_params", {}).get(
|
||||
"negative_prompt", DEFAULT_SDXL_NEGATIVE
|
||||
)
|
||||
if "---NEGATIVE---" in text:
|
||||
pos, _, neg = text.partition("---NEGATIVE---")
|
||||
pos = pos.strip()
|
||||
neg = neg.strip()
|
||||
return PreparedImagePrompt(
|
||||
prompt=_collapse_ws(pos),
|
||||
negative_prompt=neg or neg_fallback,
|
||||
)
|
||||
if "|||NEG|||" in text:
|
||||
pos, _, neg = text.partition("|||NEG|||")
|
||||
pos = pos.strip()
|
||||
neg = neg.strip()
|
||||
return PreparedImagePrompt(
|
||||
prompt=_collapse_ws(pos),
|
||||
negative_prompt=neg or neg_fallback,
|
||||
)
|
||||
return PreparedImagePrompt(
|
||||
prompt=_collapse_ws(text),
|
||||
negative_prompt=neg_fallback,
|
||||
)
|
||||
|
||||
|
||||
def _replicate_ip_adapter_scene(
|
||||
text: str, _model_config: dict[str, Any], has_reference_images: bool
|
||||
) -> PreparedImagePrompt:
|
||||
"""IP-Adapter / InstantStyle / Kolors:参考图承担风格与纹理,prompt 侧重场景与内容语义。
|
||||
|
||||
Replicate 各模型 README 均强调 prompt + 参考图配合;无参考图时由上层拦截。
|
||||
有参考图时不额外堆叠长前缀,避免稀释主体描述。
|
||||
"""
|
||||
_ = has_reference_images
|
||||
return PreparedImagePrompt(prompt=_collapse_ws(text))
|
||||
|
||||
|
||||
def _collapse_ws(s: str) -> str:
|
||||
return " ".join(s.split())
|
||||
|
||||
|
||||
def _strip_pasted_negative_block(text: str) -> str:
|
||||
lower = text.lower()
|
||||
for sep in ("\n---negative---\n", "\nnegative prompt:", "\nnegative:"):
|
||||
idx = lower.find(sep)
|
||||
if idx != -1:
|
||||
return text[:idx].strip()
|
||||
return text
|
||||
Reference in New Issue
Block a user