122 lines
4.7 KiB
Python
122 lines
4.7 KiB
Python
"""Agent 可调用的工具定义和实现。"""
|
||
|
||
from app.services.image_gen import generate_images
|
||
from app.services.view_transform import transform_view
|
||
|
||
# OpenAI Function Calling 格式的工具定义
|
||
TOOL_DEFINITIONS = [
|
||
{
|
||
"type": "function",
|
||
"function": {
|
||
"name": "generate_image",
|
||
"description": (
|
||
"根据文字描述生成图片。prompt 必须是英文。"
|
||
"如果用户提供了参考图,会自动传入 ref_image_url 参数。"
|
||
"若当前生图模型为 Stable Diffusion XL,可在正提示后单独一行写 ---NEGATIVE--- 再写负向提示;"
|
||
"不传则服务端会使用该模型的默认负向词。"
|
||
),
|
||
"parameters": {
|
||
"type": "object",
|
||
"properties": {
|
||
"prompt": {
|
||
"type": "string",
|
||
"description": "英文图片描述 prompt,详细描述要生成的图片内容、风格、颜色等",
|
||
},
|
||
"num_images": {
|
||
"type": "integer",
|
||
"description": "生成图片数量,1-4 张",
|
||
"default": 1,
|
||
"minimum": 1,
|
||
"maximum": 4,
|
||
},
|
||
},
|
||
"required": ["prompt"],
|
||
},
|
||
},
|
||
},
|
||
{
|
||
"type": "function",
|
||
"function": {
|
||
"name": "transform_view",
|
||
"description": (
|
||
"将一张图片转换为多个不同视角。需要用户先上传参考图。\n"
|
||
"可选模型:\n"
|
||
"- zero123plus(默认):~15s/次,Zero123++ 直接输出 6 个固定视角"
|
||
"(方位角 30/90/150/210/270/330°,仰角交替 30/-20°),适合快速预览,"
|
||
"对建筑大角度可能有错位。\n"
|
||
"- trellis:~105s/次,先 3D 重建再重绘,"
|
||
"支持任意方位角 + 原画风保留,适合卡通描边风格建筑。"
|
||
),
|
||
"parameters": {
|
||
"type": "object",
|
||
"properties": {
|
||
"model_id": {
|
||
"type": "string",
|
||
"enum": ["zero123plus", "trellis"],
|
||
"default": "zero123plus",
|
||
"description": "视角变换模型;默认 zero123plus(快),trellis 画风还原更强但慢",
|
||
},
|
||
"azimuths": {
|
||
"type": "array",
|
||
"items": {"type": "integer"},
|
||
"description": "mesh 管道专用:自定义方位角数组(0-359),默认 [0,60,120,180,240,300]",
|
||
},
|
||
"preserve_style": {
|
||
"type": "boolean",
|
||
"default": True,
|
||
"description": "mesh 管道专用:是否对抽帧结果走 IP-Adapter+ControlNet 重绘还原原画风",
|
||
},
|
||
},
|
||
"required": [],
|
||
},
|
||
},
|
||
},
|
||
]
|
||
|
||
|
||
async def execute_tool(
|
||
tool_name: str,
|
||
arguments: dict,
|
||
ref_image_urls: list[str] | None = None,
|
||
image_model: str | None = None,
|
||
) -> dict:
|
||
"""执行工具调用,返回结果。"""
|
||
if tool_name == "generate_image":
|
||
prompt = arguments["prompt"]
|
||
num_images = arguments.get("num_images", 1)
|
||
result = await generate_images(
|
||
prompt=prompt,
|
||
num_images=num_images,
|
||
ref_image_urls=ref_image_urls,
|
||
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": len(valid_urls) > 0,
|
||
"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,
|
||
}
|
||
|
||
if tool_name == "transform_view":
|
||
if not ref_image_urls:
|
||
return {
|
||
"success": False,
|
||
"error": "视角变换需要一张输入图片,请先上传参考图",
|
||
}
|
||
result = await transform_view(
|
||
ref_image_urls[0],
|
||
model_id=arguments.get("model_id"),
|
||
azimuths=arguments.get("azimuths"),
|
||
elevations=arguments.get("elevations"),
|
||
preserve_style=arguments.get("preserve_style", True),
|
||
)
|
||
return result
|
||
|
||
return {"success": False, "error": f"未知工具: {tool_name}"}
|