kit初版,模型引入,agent优化

This commit is contained in:
2026-04-13 00:41:23 +08:00
parent 9b053e302b
commit c435ab15cf
14097 changed files with 5032 additions and 2676248 deletions

View File

@@ -1,15 +1,19 @@
"""Agent Loop 主循环LLM 对话 -> 工具调用 -> 结果回传 -> 继续。"""
import json
import os
from typing import AsyncGenerator, Optional
from openai import AsyncOpenAI
from app.agent.tools import TOOL_DEFINITIONS, execute_tool
from app.config import get_llm_model, get_llm_max_iterations
from app.services.image_gen import to_data_uri
def _get_client() -> AsyncOpenAI:
return AsyncOpenAI()
base_url = os.getenv("OPENAI_BASE_URL")
return AsyncOpenAI(base_url=base_url) if base_url else AsyncOpenAI()
SYSTEM_PROMPT = """\
你是一个专业的游戏美术 AI 助手。你的工作是帮助美术人员通过对话生成游戏美术资源。
@@ -19,6 +23,7 @@ SYSTEM_PROMPT = """\
- 理解用户的审美意图,将中文描述转化为高质量的英文生成 prompt
- 根据用户反馈迭代修改(调整颜色、风格、构图等)
- 如果用户提供了参考图,将参考图的风格元素融入生成 prompt
- 理解用户在图片上的标注(框选区域 + 文字批注),精准定位需要修改的部分
## 工作流程
1. 理解用户需求,必要时追问细节(尺寸、风格、用途等)
@@ -26,11 +31,20 @@ SYSTEM_PROMPT = """\
3. 向用户展示结果并询问反馈
4. 根据反馈调整 prompt 并重新生成
## 标注理解
当用户发送带有「图片标注」的消息时,表示用户在之前生成的图片上做了标注。
标注格式为:`[区域 (x%, y%) 大小 w%×h%]: 修改意见`
- 区域坐标表示标注框在图片上的相对位置
- 你需要理解标注区域所指的图片内容,并将修改意见融入新的 prompt
- 如果附带了标注截图(参考图),仔细观察红色标注框和文字来理解用户意图
- 在调整 prompt 时,保持原图整体风格不变,只针对标注区域做修改
## 生成 prompt 要求
- 必须使用英文
- 尽量详细描述:主体内容、风格、颜色方案、光照、构图、材质等
- 尽量详细描述:主体内容、颜色方案、光照、构图、材质等
- 如果用户要求游戏 UI 元素,添加相关关键词如 "game UI", "icon", "button"
- 如果用户提供了参考图,在 prompt 中描述参考图的风格特征
- 如果你能直接看到参考图(图片内容),可以在 prompt 中描述参考图的风格特征
- 如果你无法看到参考图(只收到了文字提示说有参考图),参考图会由生图工具的 IP-Adapter 自动处理风格融合。此时你的 prompt 中**绝对不要猜测或指定任何画风/艺术风格关键词**(如 pixel art、watercolor、oil painting 等),只描述画面内容(主体、构图、光照等),把风格完全交给参考图来决定
## 注意事项
- 用中文和用户交流
@@ -42,6 +56,7 @@ SYSTEM_PROMPT = """\
async def run_agent_loop(
messages: list[dict],
ref_image_url: Optional[str] = None,
image_model: Optional[str] = None,
) -> AsyncGenerator[dict, None]:
"""
运行 Agent Loop以 SSE 事件流形式 yield 结果。
@@ -56,33 +71,45 @@ async def run_agent_loop(
# 构建消息列表
api_messages = [{"role": "system", "content": SYSTEM_PROMPT}]
# 检测当前 LLM 是否支持 vision多模态图片输入
llm_model = get_llm_model().lower()
vision_capable = any(kw in llm_model for kw in ("gpt-4o", "gpt-4-vision", "claude"))
for msg in messages:
if msg["role"] == "user" and ref_image_url:
# 最后一条用户消息附加参考图GPT-4o vision
if msg == messages[-1] or (
msg.get("role") == "user"
and messages.index(msg) == len(messages) - 1
):
if msg["role"] == "user" and ref_image_url and msg is messages[-1]:
if vision_capable:
# 支持 vision 的模型:直接发送图片
api_messages.append({
"role": "user",
"content": [
{"type": "text", "text": msg["content"]},
{
"type": "image_url",
"image_url": {"url": ref_image_url},
"image_url": {"url": to_data_uri(ref_image_url)},
},
],
})
continue
else:
# 不支持 vision 的模型:以文字提示告知有参考图,风格由 IP-Adapter 处理
hint = (
f"{msg['content']}\n\n"
"【系统提示:用户上传了一张参考图,已自动传递给图片生成工具的 IP-Adapter。"
"IP-Adapter 会从参考图中提取风格并融合到生成结果中。"
"你无法看到这张参考图,因此在生成 prompt 时:\n"
"1. 只描述画面内容(主体、构图、光照、材质等)\n"
"2. 不要猜测或添加任何画风/艺术风格关键词(如 pixel art、watercolor、cartoon 等)\n"
"3. 风格完全由参考图通过 IP-Adapter 决定】"
)
api_messages.append({"role": "user", "content": hint})
continue
api_messages.append({"role": msg["role"], "content": msg["content"]})
client = _get_client()
max_iterations = 5
for _ in range(max_iterations):
for _ in range(get_llm_max_iterations()):
try:
response = await client.chat.completions.create(
model="gpt-4o-mini",
model=get_llm_model(),
messages=api_messages,
tools=TOOL_DEFINITIONS,
stream=True,
@@ -163,13 +190,28 @@ async def run_agent_loop(
except json.JSONDecodeError:
arguments = {}
result = await execute_tool(tool_name, arguments, ref_image_url)
result = await execute_tool(tool_name, arguments, ref_image_url, image_model)
used_model = result.get("model_name", "")
if result.get("errors"):
yield {
"type": "tool_error",
"data": {
"tool": tool_name,
"errors": result["errors"],
"model_name": used_model,
},
}
# 如果有图片结果,推送给前端
if result.get("images"):
yield {
"type": "image_result",
"data": {"images": result["images"]},
"data": {
"images": result["images"],
"prompt_used": result.get("prompt_used", ""),
"model_name": used_model,
},
}
# 工具结果回传给 LLM
@@ -179,7 +221,7 @@ async def run_agent_loop(
"content": json.dumps(result, ensure_ascii=False),
})
# 清除 ref_image_url,避免后续轮次重复附加
ref_image_url = None
# ref_image_url 不清除:消息构建只在循环外执行一次,
# 后续工具调用仍需参考图InstantStyle 等模型必须有 style_image
yield {"type": "done", "data": {}}

View File

@@ -38,20 +38,27 @@ async def execute_tool(
tool_name: str,
arguments: dict,
ref_image_url: str | None = None,
image_model: str | None = None,
) -> dict:
"""执行工具调用,返回结果。"""
if tool_name == "generate_image":
prompt = arguments["prompt"]
num_images = arguments.get("num_images", 1)
image_urls = await generate_images(
result = await generate_images(
prompt=prompt,
num_images=num_images,
ref_image_url=ref_image_url,
model_id=image_model,
)
valid_urls = [u for u in result.urls if not u.startswith("[")]
errors = [u for u in result.urls if u.startswith("[")]
return {
"success": True,
"images": image_urls,
"success": len(valid_urls) > 0,
"images": valid_urls,
"errors": errors,
"prompt_used": prompt,
"model_name": result.model_name,
"model_id": result.model_id,
}
return {"success": False, "error": f"未知工具: {tool_name}"}

View File

@@ -7,6 +7,7 @@ from fastapi import APIRouter, File, Form, UploadFile
from sse_starlette.sse import EventSourceResponse
from app.agent.loop import run_agent_loop
from app.config import get_image_models_list, get_default_image_model_id
router = APIRouter()
@@ -23,26 +24,56 @@ async def _save_upload(file: UploadFile) -> str:
return f"/uploads/{filename}"
@router.post("/upload-ref-image")
async def upload_ref_image(
file: UploadFile = File(...),
):
"""
独立的参考图上传端点。
前端选图后立即调用,返回服务端路径,供后续发消息时引用。
"""
UPLOADS_DIR.mkdir(parents=True, exist_ok=True)
url = await _save_upload(file)
return {"url": url, "filename": file.filename}
@router.get("/models")
async def list_models():
"""返回可用的图像生成模型列表。"""
return {
"models": get_image_models_list(),
"default": get_default_image_model_id(),
}
@router.post("/chat")
async def chat(
messages: str = Form(...),
ref_image: Optional[UploadFile] = File(None),
ref_image_url: Optional[str] = Form(None),
image_model: Optional[str] = Form(None),
):
"""
主对话端点。
参数:
- messages: JSON 字符串,对话历史 [{role, content}]
- ref_image: 可选的参考图文件
- ref_image: 可选的参考图文件(兼容旧方式)
- ref_image_url: 可选,已通过 /upload-ref-image 上传后的服务端路径
- image_model: 可选,指定本次使用的生图模型短 ID
"""
parsed_messages = json.loads(messages)
ref_image_url: Optional[str] = None
if ref_image and ref_image.filename:
ref_image_url = await _save_upload(ref_image)
resolved_ref_url: Optional[str] = None
if ref_image_url:
resolved_ref_url = ref_image_url
elif ref_image and ref_image.filename:
resolved_ref_url = await _save_upload(ref_image)
async def event_generator():
async for event in run_agent_loop(parsed_messages, ref_image_url):
async for event in run_agent_loop(
parsed_messages, resolved_ref_url, image_model=image_model
):
yield {
"event": event["type"],
"data": json.dumps(event["data"], ensure_ascii=False),

View File

@@ -0,0 +1,156 @@
"""
EPEEKit 集中配置。
所有可调参数从环境变量读取,每次调用实时读取(不缓存),确保 load_dotenv() 后生效。
"""
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"))
# ─── 图像模型注册表 ─────────────────────────────────────
#
# 每个模型的配置说明:
# id — 前端/API 使用的短 ID
# name — 显示名称
# provider — 生成服务提供者(对应 image_gen.py 中的 Provider
# model_id — Replicate 上的完整模型 ID
# description — 前端下拉列表中的说明文字
# supports_ref_image — 是否原生支持参考图输入IP-Adapter 等)
# ref_image_param — 传给 Replicate 的参考图参数名(模型间可能不同)
# num_images_param — 批量生成参数名Flux 用 num_outputsKolors 用 number_of_images
# default_params — 默认推理参数
IMAGE_MODELS: dict[str, dict[str, Any]] = {
"flux-schnell": {
"id": "flux-schnell",
"name": "Flux Schnell",
"provider": "replicate",
"model_id": "black-forest-labs/flux-schnell",
"description": "快速生成,适合快速迭代",
"supports_ref_image": False,
"num_images_param": "num_outputs",
"default_params": {
"aspect_ratio": "1:1",
"output_format": "png",
},
},
"flux-dev": {
"id": "flux-dev",
"name": "Flux Dev",
"provider": "replicate",
"model_id": "black-forest-labs/flux-dev",
"description": "高质量生成,细节更好",
"supports_ref_image": False,
"num_images_param": "num_outputs",
"default_params": {
"aspect_ratio": "1:1",
"output_format": "png",
},
},
"sdxl": {
"id": "sdxl",
"name": "Stable Diffusion XL",
"provider": "replicate",
"model_id": "stability-ai/sdxl",
"description": "经典 SDXL支持 negative prompt",
"supports_ref_image": False,
"num_images_param": "num_outputs",
"default_params": {
"width": 1024,
"height": 1024,
"num_inference_steps": 50,
"guidance_scale": 7.5,
},
},
"instant-style": {
"id": "instant-style",
"name": "InstantStyle",
"provider": "replicate",
"model_id": "jyoung105/instant-style:c6f01e12f31cb99f9ee774a78992a71294f630a6f433d9aecfdc33b816fc4baa",
"description": "强风格迁移,画风还原度高(较慢)",
"supports_ref_image": True,
"ref_image_param": "style_image",
"num_images_param": "num_outputs",
"default_params": {
"width": 1024,
"height": 1024,
"num_inference_steps": 30,
"guidance_scale": 5,
"style_strength": 1.0,
"block_mode": "style-only",
"adapter_mode": "original",
},
},
"kolors-ipadapter": {
"id": "kolors-ipadapter",
"name": "Kolors IP-Adapter",
"provider": "replicate",
"model_id": "fofr/kolors-with-ipadapter:5a1a92b2c0f81813225d48ed8e411813da41aa84e7582fb705d1af46eea36eed",
"description": "风格参考生成,上传参考图效果最佳",
"supports_ref_image": True,
"ref_image_param": "image",
"num_images_param": "number_of_images",
"default_params": {
"width": 1024,
"height": 1024,
"steps": 25,
"cfg": 4,
"ip_adapter_weight": 0.8,
"ip_adapter_weight_type": "style transfer precise",
"output_format": "png",
},
},
}
def get_default_image_model_id() -> str:
"""返回 .env 中配置的默认模型短 ID若不在注册表中则回退到 flux-schnell。"""
env_model = os.getenv("IMAGE_MODEL", "flux-schnell")
for mid, cfg in IMAGE_MODELS.items():
if cfg["model_id"] == env_model or mid == env_model:
return mid
return "flux-schnell"
def get_image_model_config(model_id: str | None = None) -> dict[str, Any]:
"""根据短 ID 获取模型配置,未指定或不存在则使用默认模型。"""
if model_id and model_id in IMAGE_MODELS:
return IMAGE_MODELS[model_id]
return IMAGE_MODELS[get_default_image_model_id()]
def get_ref_image_model_id() -> str | None:
"""返回有参考图时推荐使用的模型 ID第一个 supports_ref_image=True 的模型)。"""
for mid, cfg in IMAGE_MODELS.items():
if cfg.get("supports_ref_image"):
return mid
return None
def get_image_models_list() -> list[dict]:
"""返回前端下拉列表所需的模型摘要信息。"""
return [
{
"id": cfg["id"],
"name": cfg["name"],
"description": cfg["description"],
"supports_ref_image": cfg.get("supports_ref_image", False),
}
for cfg in IMAGE_MODELS.values()
]
def get_image_aspect_ratio() -> str:
return os.getenv("IMAGE_ASPECT_RATIO", "1:1")
def get_image_output_format() -> str:
return os.getenv("IMAGE_OUTPUT_FORMAT", "png")

View File

@@ -2,15 +2,15 @@ import os
from pathlib import Path
from dotenv import load_dotenv
load_dotenv(override=True)
from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware
from fastapi.staticfiles import StaticFiles
from app.api.chat import router as chat_router
load_dotenv()
app = FastAPI(title="Art Agent MVP")
app = FastAPI(title="EPEEKit API")
app.add_middleware(
CORSMiddleware,

View File

@@ -1,76 +1,220 @@
"""Replicate 图像生成服务封装"""
"""图像生成服务 — Provider 抽象层"""
import base64
import mimetypes
import os
import uuid
from abc import ABC, abstractmethod
from pathlib import Path
from typing import Any
import httpx
import replicate
from replicate import Client as ReplicateClient
GENERATED_DIR = Path(__file__).parent.parent.parent / "generated"
from app.config import get_image_model_config, get_image_aspect_ratio, get_image_output_format
BACKEND_ROOT = Path(__file__).parent.parent.parent
GENERATED_DIR = BACKEND_ROOT / "generated"
def _make_replicate_client() -> ReplicateClient:
"""
创建 Replicate 客户端。
Replicate SDK 内部创建 httpx transport 时不读取代理环境变量,
需要我们手动把代理配置注入到 transport 中。
"""
timeout = httpx.Timeout(5.0, read=300.0, write=30.0, connect=30.0, pool=10.0)
proxy_url = os.environ.get("HTTPS_PROXY") or os.environ.get("HTTP_PROXY")
transport_kwargs: dict[str, Any] = {}
if proxy_url:
transport_kwargs["proxy"] = proxy_url
our_transport = httpx.AsyncHTTPTransport(**transport_kwargs)
client = ReplicateClient(
timeout=timeout,
transport=our_transport,
)
return client
_replicate_client = _make_replicate_client()
# ─── 公共工具 ──────────────────────────────────────────
def to_data_uri(image_path: str) -> str:
"""将本地路径或已有 URL 转为 Replicate 可接受的格式data URI 或原始 URL"""
if image_path.startswith("data:") or image_path.startswith("http"):
return image_path
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 f"data:{mime};base64,{b64}"
return image_path
async def _download_image(url: str) -> str:
"""下载远程图片到本地 generated/ 目录,返回本地 URL 路径。"""
filename = f"{uuid.uuid4().hex}.png"
filepath = GENERATED_DIR / filename
async with httpx.AsyncClient() as client:
proxy = os.environ.get("HTTPS_PROXY") or os.environ.get("HTTP_PROXY")
async with httpx.AsyncClient(proxy=proxy, timeout=httpx.Timeout(60.0)) as client:
resp = await client.get(url, follow_redirects=True)
resp.raise_for_status()
filepath.write_bytes(resp.content)
return f"/generated/{filename}"
# ─── Provider 抽象基类 ─────────────────────────────────
class ImageProvider(ABC):
"""所有图像生成 provider 的基类。"""
@abstractmethod
async def generate(
self,
model_config: dict[str, Any],
prompt: str,
num_images: int = 1,
ref_image_url: str | None = None,
) -> list[str]:
"""生成图片并返回本地 URL 列表。"""
...
class ReplicateProvider(ImageProvider):
"""通过 Replicate API 调用模型。"""
async def generate(
self,
model_config: dict[str, Any],
prompt: str,
num_images: int = 1,
ref_image_url: str | None = None,
) -> list[str]:
replicate_model_id = model_config["model_id"]
default_params: dict[str, Any] = model_config.get("default_params", {})
local_urls: list[str] = []
try:
input_params = self._build_input(
model_config, default_params, prompt, num_images, ref_image_url
)
output = await _replicate_client.async_run(
replicate_model_id, input=input_params, wait=False
)
items = output if isinstance(output, list) else [output]
for item in items:
url = str(item)
if url.startswith("https://") or url.startswith("http://") or url.startswith("data:"):
local_urls.append(await _download_image(url))
else:
local_urls.append(f"[生成失败: 模型返回非图片内容: {url[:200]}]")
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_input(
model_config: dict[str, Any],
default_params: dict[str, Any],
prompt: str,
num_images: int = 1,
ref_image_url: str | None = None,
) -> dict[str, Any]:
"""
根据模型配置构建 Replicate input 参数。
通过 model_config 中的标志字段自动适配不同模型:
- supports_ref_image / ref_image_param: 参考图注入
- num_images_param: 各模型的批量生成参数名(如 number_of_images / num_outputs
"""
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():
params[k] = v
params[num_images_param] = num_images
if ref_image_url:
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}
if is_flux:
params["aspect_ratio"] = default_params.get(
"aspect_ratio", get_image_aspect_ratio()
)
params["output_format"] = default_params.get(
"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"]
return params
# ─── Provider 注册 ─────────────────────────────────────
_PROVIDERS: dict[str, ImageProvider] = {
"replicate": ReplicateProvider(),
}
# ─── 公开 API ──────────────────────────────────────────
class GenerateResult:
"""图片生成结果,包含生成的 URL 列表和实际使用的模型信息。"""
def __init__(self, urls: list[str], model_name: str, model_id: str):
self.urls = urls
self.model_name = model_name
self.model_id = model_id
async def generate_images(
prompt: str,
num_images: int = 1,
ref_image_url: str | None = None,
) -> list[str]:
model_id: str | None = None,
) -> GenerateResult:
"""
调用 Replicate 生成图片
统一入口:根据 model_id 查注册表,分发到对应 provider
返回本地可访问的图片 URL 列表
model_id 为空时使用 .env 中配置的默认模型
"""
local_urls = []
config = get_image_model_config(model_id)
provider_name = config.get("provider", "replicate")
provider = _PROVIDERS.get(provider_name)
for _ in range(num_images):
try:
if ref_image_url and ref_image_url.startswith("/"):
# 本地路径转为 file URI 不适用于 Replicate
# 需要用户上传的图先通过后端 URL 访问
# MVP 阶段:参考图作为 prompt 的文字补充,不直接传给模型
# 仅使用 flux-schnell 文生图
output = await replicate.async_run(
"black-forest-labs/flux-schnell",
input={
"prompt": prompt,
"num_outputs": 1,
"aspect_ratio": "1:1",
"output_format": "png",
},
)
else:
output = await replicate.async_run(
"black-forest-labs/flux-schnell",
input={
"prompt": prompt,
"num_outputs": 1,
"aspect_ratio": "1:1",
"output_format": "png",
},
)
model_name = config.get("name", "未知模型")
resolved_id = config.get("id", model_id or "unknown")
# output 是 FileOutput 列表或单个 URL
if isinstance(output, list):
for item in output:
url = str(item)
local_url = await _download_image(url)
local_urls.append(local_url)
else:
url = str(output)
local_url = await _download_image(url)
local_urls.append(local_url)
if not provider:
return GenerateResult([f"[未知 provider: {provider_name}]"], model_name, resolved_id)
except Exception as e:
local_urls.append(f"[生成失败: {e}]")
return local_urls
urls = await provider.generate(config, prompt, num_images, ref_image_url)
return GenerateResult(urls, model_name, resolved_id)