1365 lines
58 KiB
Python
1365 lines
58 KiB
Python
"""
|
||
Comfyui_o1key - ComfyUI 自定义节点集合
|
||
通过 api.o1key.cn 调用 AI 模型进行图像生成和文本生成
|
||
|
||
项目结构:
|
||
├── nodes/ # 节点实现
|
||
├── utils/ # 工具模块
|
||
├── clients/ # API 客户端
|
||
└── __init__.py # 节点注册入口
|
||
"""
|
||
|
||
|
||
import ssl
|
||
import logging
|
||
import asyncio
|
||
|
||
# 屏蔽 ComfyUI 资产扫描的终端日志输出
|
||
_seeder_filter = lambda record: not any(
|
||
kw in record.getMessage()
|
||
for kw in ("Seeder start", "Asset scan", "Scan(", "Fast scan")
|
||
)
|
||
logging.getLogger().addFilter(_seeder_filter)
|
||
|
||
|
||
def _is_ignored_asyncio_win10054(context):
|
||
exc = context.get("exception")
|
||
if not (
|
||
isinstance(exc, ConnectionResetError)
|
||
and getattr(exc, "winerror", None) == 10054
|
||
):
|
||
return False
|
||
|
||
handle = str(context.get("handle", ""))
|
||
message = str(context.get("message", ""))
|
||
marker = "_ProactorBasePipeTransport._call_connection_lost"
|
||
return marker in handle or marker in message
|
||
|
||
|
||
def _install_asyncio_win10054_filter(loop):
|
||
if getattr(loop, "_o1key_win10054_filter_installed", False):
|
||
return loop
|
||
|
||
previous_handler = loop.get_exception_handler()
|
||
|
||
def _o1key_asyncio_exception_handler(loop, context):
|
||
if _is_ignored_asyncio_win10054(context):
|
||
return
|
||
if previous_handler is not None:
|
||
previous_handler(loop, context)
|
||
else:
|
||
loop.default_exception_handler(context)
|
||
|
||
loop.set_exception_handler(_o1key_asyncio_exception_handler)
|
||
setattr(loop, "_o1key_win10054_filter_installed", True)
|
||
return loop
|
||
|
||
|
||
try:
|
||
_install_asyncio_win10054_filter(asyncio.get_event_loop())
|
||
except RuntimeError:
|
||
pass
|
||
|
||
if not getattr(asyncio, "_o1key_new_event_loop_patched", False):
|
||
_o1key_original_new_event_loop = asyncio.new_event_loop
|
||
|
||
def _o1key_new_event_loop(*args, **kwargs):
|
||
return _install_asyncio_win10054_filter(
|
||
_o1key_original_new_event_loop(*args, **kwargs)
|
||
)
|
||
|
||
asyncio.new_event_loop = _o1key_new_event_loop
|
||
asyncio._o1key_new_event_loop_patched = True
|
||
|
||
from .nodes import NanoBananaPro, BatchNanoBananaPro, GoogleGemini, LoadFile, LoadImagesFromFolder, ImageStitchPro, BatchCleanMetadata, VideoPreview, GoogleVeo, Google31Video, MiniMaxH3Video, FluxImageEdit, UniversalLLMChat, BatchImagesO1key, SeedanceMultiModal, StreamPreview, DoubaoImage, O1keyGPTImage, O1keyGPTImageBatch, O1keyGrokImage, O1keyGrokVideo, O1keyGrokVideoEdit
|
||
from .nodes import K3Video, K3MotionControl, SaveImageFormat
|
||
from .nodes import O1keySavePSD
|
||
from .nodes import O1keyRemoveBackground
|
||
from .nodes import O1keyGridSplitter
|
||
from .nodes import O1keyPromptMultiFunction
|
||
from .nodes import O1keyVideoTrim
|
||
from .nodes import SeedanceElementCreate
|
||
from .nodes import SeedanceAutoPass
|
||
from .nodes import SeedanceAutoPassBatch
|
||
from .nodes import O1keyAutoRedCast
|
||
from .nodes import O1keyImageGenerator, O1keyImageSave
|
||
from .nodes import O1keyVideoGenerator, O1keyVideoResult
|
||
from .nodes import O1keyOmniFlashVideo
|
||
|
||
# 报错弹框友好文案(不修改原节点代码,仅在外层统一处理)
|
||
_MSG_TIMEOUT = "API 请求超时,请稍后重试或检查网络。"
|
||
_MSG_SSL_NETWORK = (
|
||
"本地网络不太稳定!解决方案如下:\n"
|
||
"1. 重启程序再试试看 (优先)\n"
|
||
"2. 调整一下网络环境,如wifi或宽带等\n"
|
||
"3. 切换VPN节点,或更换代理模式\n"
|
||
"4. 关掉杀毒软件或防火墙\n"
|
||
"5. 关掉浏览器VPN插件,避免冲突"
|
||
)
|
||
|
||
def _wrap_generate_for_error_display(cls, attr="generate"):
|
||
original = getattr(cls, attr, None)
|
||
if original is None:
|
||
return
|
||
def wrapped(self, *args, **kwargs):
|
||
try:
|
||
return original(self, *args, **kwargs)
|
||
except TimeoutError as e:
|
||
msg = (str(e) or "").strip()
|
||
if not msg:
|
||
msg = _MSG_TIMEOUT
|
||
raise TimeoutError(msg) from None
|
||
except (ssl.SSLError, OSError) as e:
|
||
err_str = str(e)
|
||
if "DECRYPTION_FAILED_OR_BAD_RECORD_MAC" in err_str or "decryption failed or bad record mac" in err_str.lower():
|
||
raise RuntimeError(_MSG_SSL_NETWORK) from None
|
||
raise
|
||
setattr(cls, attr, wrapped)
|
||
|
||
_wrap_generate_for_error_display(NanoBananaPro)
|
||
_wrap_generate_for_error_display(BatchNanoBananaPro)
|
||
|
||
# ComfyUI 节点注册
|
||
NODE_CLASS_MAPPINGS = {
|
||
"NanoBanana": NanoBananaPro,
|
||
"BatchNanoBananaPro": BatchNanoBananaPro,
|
||
"GoogleGemini": GoogleGemini,
|
||
"LoadFile": LoadFile,
|
||
"O1keyLoadImagesFromFolder": LoadImagesFromFolder,
|
||
"ImageStitchPro": ImageStitchPro,
|
||
|
||
"BatchCleanMetadata": BatchCleanMetadata,
|
||
"VideoPreview": VideoPreview,
|
||
"GoogleVeo": GoogleVeo,
|
||
"Google31Video": Google31Video,
|
||
"MiniMaxH3Video": MiniMaxH3Video,
|
||
"FluxImageEdit": FluxImageEdit,
|
||
"UniversalLLMChat": UniversalLLMChat,
|
||
|
||
"BatchImagesO1key": BatchImagesO1key,
|
||
"SeedanceMultiModal": SeedanceMultiModal,
|
||
"StreamPreview": StreamPreview,
|
||
"DoubaoImage": DoubaoImage,
|
||
"O1keyGPTImage": O1keyGPTImage,
|
||
"O1keyGPTImageBatch": O1keyGPTImageBatch,
|
||
"O1keyGrokImage": O1keyGrokImage,
|
||
"O1keyGrokVideo": O1keyGrokVideo,
|
||
"O1keyGrokVideoEdit": O1keyGrokVideoEdit,
|
||
"K3Video": K3Video,
|
||
"K3MotionControl": K3MotionControl,
|
||
"SaveImageFormat": SaveImageFormat,
|
||
"O1keySavePSD": O1keySavePSD,
|
||
"O1keyRemoveBackground": O1keyRemoveBackground,
|
||
"O1keyGridSplitter": O1keyGridSplitter,
|
||
"O1keyPromptMultiFunction": O1keyPromptMultiFunction,
|
||
"O1keyVideoTrim": O1keyVideoTrim,
|
||
"SeedanceElementCreate": SeedanceElementCreate,
|
||
"SeedanceAutoPass": SeedanceAutoPass,
|
||
"SeedanceAutoPassBatch": SeedanceAutoPassBatch,
|
||
"O1keyAutoRedCast": O1keyAutoRedCast,
|
||
"O1keyImageGenerator": O1keyImageGenerator,
|
||
"O1keyImageSave": O1keyImageSave,
|
||
"O1keyVideoGenerator": O1keyVideoGenerator,
|
||
"O1keyVideoResult": O1keyVideoResult,
|
||
"O1keyOmniFlashVideo": O1keyOmniFlashVideo,
|
||
}
|
||
|
||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||
"O1keyVideoGenerator": "o1key 视频生成",
|
||
"O1keyVideoResult": "o1key 视频结果",
|
||
"O1keyOmniFlashVideo": "Omni Flash 视频生成",
|
||
"NanoBanana": "Nano Banana",
|
||
"BatchNanoBananaPro": "Nano Banana 批量跑图",
|
||
"GoogleGemini": "Google Gemini",
|
||
"LoadFile": "加载文件",
|
||
"O1keyLoadImagesFromFolder": "加载图像(文件夹)",
|
||
"ImageStitchPro": "图像拼接 Pro",
|
||
|
||
"BatchCleanMetadata": "批量任务(防AI识别)",
|
||
"VideoPreview": "预览视频",
|
||
"GoogleVeo": "Google Veo - ab",
|
||
"Google31Video": "Google 3.1 Video",
|
||
"MiniMaxH3Video": "MiniMax H3 / H3 Max 视频生成",
|
||
"FluxImageEdit": "Flux2 图像编辑",
|
||
"UniversalLLMChat": "提示词专家",
|
||
|
||
"BatchImagesO1key": "加载图像(批量)",
|
||
"SeedanceMultiModal": "Seedance 多模态参考生视频",
|
||
"StreamPreview": "流式文本预览",
|
||
"DoubaoImage": "豆包生图",
|
||
"O1keyGPTImage": "gpt image",
|
||
"O1keyGPTImageBatch": "o1key GPT Image(批量)",
|
||
"O1keyGrokImage": "Grok Image",
|
||
"O1keyGrokVideo": "Grok Video",
|
||
"O1keyGrokVideoEdit": "Grok Video Edit",
|
||
"K3Video": "K 视频生成",
|
||
"K3MotionControl": "K 动作模仿",
|
||
"SaveImageFormat": "保存图像(格式转换)",
|
||
"O1keySavePSD": "保存 PSD(分层)",
|
||
"O1keyRemoveBackground": "去背景(rembg)",
|
||
"O1keyGridSplitter": "合并图智能切割",
|
||
"O1keyPromptMultiFunction": "提示词(多功能)",
|
||
"O1keyVideoTrim": "视频裁剪",
|
||
"SeedanceElementCreate": "Seedance 创建素材",
|
||
"SeedanceAutoPass": "Seedance 全能生成视频",
|
||
"SeedanceAutoPassBatch": "Seedance 全能生成视频(批量)",
|
||
"O1keyAutoRedCast": "自动红偏校正",
|
||
"O1keyImageGenerator": "o1key 图片生成",
|
||
"O1keyImageSave": "o1key 保存图像",
|
||
}
|
||
|
||
WEB_DIRECTORY = "./web"
|
||
|
||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS', 'WEB_DIRECTORY']
|
||
|
||
# 注册 /o1key/input_dir 接口,供前端文件上传按钮获取 input 目录绝对路径
|
||
try:
|
||
from aiohttp import web
|
||
from server import PromptServer
|
||
import folder_paths
|
||
from .utils.config import (
|
||
DEFAULT_NETWORK_ROUTE,
|
||
NETWORK_ROUTE_CONFIG_KEY,
|
||
NETWORK_ROUTE_OPTIONS,
|
||
NETWORK_ROUTES,
|
||
get_api_key,
|
||
get_network_route,
|
||
load_config,
|
||
update_config,
|
||
)
|
||
from .utils.chat_support import (
|
||
PROMPT_OPTIMIZER_TIMEOUT_SECONDS,
|
||
build_search_context,
|
||
expand_xlsx_attachments,
|
||
extract_search_query,
|
||
optimize_image_prompt,
|
||
rewrite_search_query,
|
||
web_search,
|
||
write_video_prompt,
|
||
)
|
||
from .utils.o1key_image_jobs import register_o1key_image_job_routes
|
||
from .utils.o1key_image_thumbnail import register_o1key_image_thumbnail_route
|
||
from .utils.o1key_video_jobs import register_o1key_video_job_routes
|
||
from .utils.updater import UpdateError, check_for_update, update_package
|
||
|
||
_O1KEY_IMAGE_JOB_MANAGER = register_o1key_image_job_routes(
|
||
PromptServer,
|
||
web,
|
||
folder_paths,
|
||
)
|
||
register_o1key_image_thumbnail_route(PromptServer, web, folder_paths)
|
||
_O1KEY_VIDEO_JOB_MANAGER = register_o1key_video_job_routes(
|
||
PromptServer,
|
||
web,
|
||
folder_paths,
|
||
)
|
||
|
||
# 每次 ComfyUI 进程启动都会生成新的标识。前端据此确认后端确实完成了
|
||
# 重启,而不是仅仅重新加载了浏览器页面。
|
||
import os as _restart_os
|
||
import sys as _restart_sys
|
||
import threading as _restart_threading
|
||
import time as _restart_time
|
||
import uuid as _restart_uuid
|
||
|
||
_O1KEY_BOOT_ID = _restart_uuid.uuid4().hex
|
||
_o1key_restart_pending = _restart_threading.Event()
|
||
_o1key_update_lock = _restart_threading.Lock()
|
||
|
||
@PromptServer.instance.routes.get("/o1key/update/check")
|
||
async def check_o1key_update(request):
|
||
if not _o1key_update_lock.acquire(blocking=False):
|
||
return web.json_response(
|
||
{"code": "update_in_progress", "error": "更新正在进行。"}, status=409,
|
||
)
|
||
try:
|
||
result = await asyncio.to_thread(check_for_update)
|
||
return web.json_response(result)
|
||
except UpdateError as exc:
|
||
return web.json_response(exc.as_dict(), status=exc.status)
|
||
except Exception:
|
||
logging.exception("O1Key 检查更新失败")
|
||
return web.json_response(
|
||
{"code": "internal_error", "error": "检查更新失败。"}, status=500,
|
||
)
|
||
finally:
|
||
_o1key_update_lock.release()
|
||
|
||
@PromptServer.instance.routes.post("/o1key/update")
|
||
async def update_o1key_package(request):
|
||
if request.headers.get("X-O1Key-Update") != "1":
|
||
return web.json_response(
|
||
{"code": "invalid_request", "error": "无效的更新请求。", "suggestion": "请从 O1Key 更新面板重新操作。"},
|
||
status=403,
|
||
)
|
||
if not _o1key_update_lock.acquire(blocking=False):
|
||
return web.json_response(
|
||
{"code": "update_in_progress", "error": "更新正在进行。", "suggestion": "请等待当前操作完成,不要重复点击。"},
|
||
status=409,
|
||
)
|
||
try:
|
||
result = await asyncio.to_thread(update_package)
|
||
return web.json_response(result)
|
||
except UpdateError as exc:
|
||
return web.json_response(exc.as_dict(), status=exc.status)
|
||
except Exception:
|
||
logging.exception("O1Key 更新失败")
|
||
return web.json_response(
|
||
{"code": "internal_error", "error": "更新失败。", "suggestion": "请查看 ComfyUI 日志,并在确认本地文件安全后重试。"},
|
||
status=500,
|
||
)
|
||
finally:
|
||
_o1key_update_lock.release()
|
||
|
||
def _o1key_restart_command():
|
||
"""复用当前解释器和启动参数,并禁止重启时额外打开浏览器。"""
|
||
auto_launch_flags = {"--auto-launch", "--auto_launch", "--launch"}
|
||
# orig_argv 包含嵌入式 Python 的 -s 等解释器参数;普通 sys.argv 不包含。
|
||
# 保留这些参数可确保便携版重启前后的运行环境完全一致。
|
||
original = getattr(_restart_sys, "orig_argv", None)
|
||
source_arguments = original[1:] if original else _restart_sys.argv
|
||
arguments = [arg for arg in source_arguments if arg not in auto_launch_flags]
|
||
if "--disable-auto-launch" not in arguments:
|
||
arguments.append("--disable-auto-launch")
|
||
return [_restart_sys.executable, *arguments]
|
||
|
||
def _restart_o1key_comfyui_process(delay=1.25):
|
||
"""在响应发送完成后,用相同终端进程重新启动 ComfyUI。"""
|
||
try:
|
||
_restart_time.sleep(delay)
|
||
try:
|
||
_restart_sys.stdout.flush()
|
||
_restart_sys.stderr.flush()
|
||
except Exception:
|
||
pass
|
||
|
||
command = _o1key_restart_command()
|
||
print("[O1Key] 正在重启 ComfyUI...", flush=True)
|
||
_restart_os.execv(command[0], command)
|
||
except Exception:
|
||
_o1key_restart_pending.clear()
|
||
logging.exception("O1Key 无法重启 ComfyUI 进程")
|
||
|
||
@PromptServer.instance.routes.get("/o1key/restart/status")
|
||
async def get_o1key_restart_status(request):
|
||
return web.json_response(
|
||
{
|
||
"ready": True,
|
||
"boot_id": _O1KEY_BOOT_ID,
|
||
"pid": _restart_os.getpid(),
|
||
},
|
||
headers={"Cache-Control": "no-store"},
|
||
)
|
||
|
||
@PromptServer.instance.routes.post("/o1key/restart")
|
||
async def restart_o1key_comfyui(request):
|
||
if _o1key_restart_pending.is_set():
|
||
return web.json_response(
|
||
{
|
||
"success": False,
|
||
"error": "ComfyUI 正在重启,请稍候。",
|
||
"boot_id": _O1KEY_BOOT_ID,
|
||
},
|
||
status=409,
|
||
)
|
||
|
||
_o1key_restart_pending.set()
|
||
worker = _restart_threading.Thread(
|
||
target=_restart_o1key_comfyui_process,
|
||
name="o1key-comfyui-restart",
|
||
daemon=True,
|
||
)
|
||
worker.start()
|
||
return web.json_response(
|
||
{
|
||
"success": True,
|
||
"message": "ComfyUI 正在重启。",
|
||
"boot_id": _O1KEY_BOOT_ID,
|
||
"pid": _restart_os.getpid(),
|
||
},
|
||
headers={"Cache-Control": "no-store"},
|
||
)
|
||
|
||
def _get_o1key_server_port():
|
||
try:
|
||
import comfy.cli_args as _cli_args
|
||
args = getattr(_cli_args, "args", None)
|
||
port = getattr(args, "port", None) if args else None
|
||
port = port or getattr(_cli_args, "server_port", None) or getattr(_cli_args, "port", None)
|
||
if port is not None:
|
||
return str(int(port))
|
||
except Exception:
|
||
pass
|
||
try:
|
||
import sys as _sys
|
||
for idx, arg in enumerate(_sys.argv):
|
||
if arg in ("--port", "--listen-port") and idx + 1 < len(_sys.argv):
|
||
return str(int(_sys.argv[idx + 1]))
|
||
for prefix in ("--port=", "--listen-port="):
|
||
if arg.startswith(prefix):
|
||
return str(int(arg.split("=", 1)[1]))
|
||
except Exception:
|
||
pass
|
||
return "8188"
|
||
|
||
def _get_o1key_history_meta_file(output_dir):
|
||
import os as _os_history
|
||
return _os_history.path.join(
|
||
output_dir,
|
||
f".o1key_history_{_get_o1key_server_port()}.json",
|
||
)
|
||
|
||
def _get_o1key_notes_file():
|
||
# 笔记固定存 ComfyUI input 目录,插件更新/替换不会清空笔记
|
||
import os as _os_notes
|
||
input_dir = _os_notes.path.abspath(folder_paths.get_input_directory())
|
||
_os_notes.makedirs(input_dir, exist_ok=True)
|
||
notes_file = _os_notes.path.join(input_dir, "o1key-notes.json")
|
||
|
||
return notes_file
|
||
|
||
def _get_o1key_cases_dir():
|
||
import os as _os_cases
|
||
cases_dir = _os_cases.path.join(_os_cases.path.dirname(__file__), "cases")
|
||
_os_cases.makedirs(cases_dir, exist_ok=True)
|
||
return cases_dir
|
||
|
||
def _get_o1key_case_file(filename):
|
||
import os as _os_cases
|
||
safe_name = _os_cases.path.basename(filename or "")
|
||
if not safe_name.lower().endswith(".json"):
|
||
return None
|
||
cases_dir = _get_o1key_cases_dir()
|
||
path = _os_cases.path.abspath(_os_cases.path.join(cases_dir, safe_name))
|
||
if not path.startswith(_os_cases.path.abspath(cases_dir) + _os_cases.sep):
|
||
return None
|
||
return path
|
||
|
||
def _extract_o1key_notes(payload):
|
||
if isinstance(payload, list):
|
||
return payload
|
||
if isinstance(payload, dict) and isinstance(payload.get("notes"), list):
|
||
return payload["notes"]
|
||
return None
|
||
|
||
@PromptServer.instance.routes.get("/o1key/cases")
|
||
async def get_o1key_cases(request):
|
||
import os as _os_cases
|
||
import json as _json_cases
|
||
|
||
cases_dir = _get_o1key_cases_dir()
|
||
cases = []
|
||
for filename in sorted(_os_cases.listdir(cases_dir), key=str.lower):
|
||
if not filename.lower().endswith(".json"):
|
||
continue
|
||
path = _get_o1key_case_file(filename)
|
||
if not path or not _os_cases.path.isfile(path):
|
||
continue
|
||
title = _os_cases.path.splitext(filename)[0]
|
||
try:
|
||
with open(path, "r", encoding="utf-8") as cf:
|
||
data = _json_cases.load(cf)
|
||
if isinstance(data, dict):
|
||
title = str(data.get("title") or data.get("name") or title)
|
||
except Exception:
|
||
pass
|
||
cases.append({"id": filename, "filename": filename, "title": title})
|
||
return web.json_response({"cases": cases, "path": str(cases_dir)})
|
||
|
||
@PromptServer.instance.routes.get("/o1key/case")
|
||
async def get_o1key_case(request):
|
||
import os as _os_cases
|
||
import json as _json_cases
|
||
|
||
filename = request.query.get("file", "")
|
||
path = _get_o1key_case_file(filename)
|
||
if not path or not _os_cases.path.isfile(path):
|
||
return web.json_response({"error": "case not found"}, status=404)
|
||
try:
|
||
with open(path, "r", encoding="utf-8") as cf:
|
||
data = _json_cases.load(cf)
|
||
except Exception as e:
|
||
return web.json_response({"error": str(e)}, status=500)
|
||
return web.json_response({"filename": _os_cases.path.basename(path), "case": data})
|
||
|
||
@PromptServer.instance.routes.get("/o1key/notes")
|
||
async def get_o1key_notes(request):
|
||
import os as _os_notes
|
||
import json as _json_notes
|
||
|
||
notes_file = _get_o1key_notes_file()
|
||
exists = _os_notes.path.isfile(notes_file)
|
||
notes = []
|
||
|
||
if exists:
|
||
try:
|
||
with open(notes_file, "r", encoding="utf-8") as nf:
|
||
loaded = _json_notes.load(nf)
|
||
notes = _extract_o1key_notes(loaded)
|
||
if notes is None:
|
||
return web.json_response(
|
||
{"error": "invalid notes file", "path": notes_file},
|
||
status=500,
|
||
)
|
||
except Exception as e:
|
||
return web.json_response(
|
||
{"error": str(e), "path": notes_file},
|
||
status=500,
|
||
)
|
||
|
||
return web.json_response({"notes": notes, "path": notes_file, "exists": exists})
|
||
|
||
@PromptServer.instance.routes.post("/o1key/notes")
|
||
async def save_o1key_notes(request):
|
||
import os as _os_notes
|
||
import json as _json_notes
|
||
|
||
try:
|
||
payload = await request.json()
|
||
notes = _extract_o1key_notes(payload)
|
||
if notes is None:
|
||
return web.json_response({"error": "notes must be a list"}, status=400)
|
||
except Exception as e:
|
||
return web.json_response({"error": f"invalid notes payload: {str(e)}"}, status=400)
|
||
|
||
notes_file = _get_o1key_notes_file()
|
||
temp_file = notes_file + ".tmp"
|
||
try:
|
||
with open(temp_file, "w", encoding="utf-8") as nf:
|
||
_json_notes.dump(notes, nf, ensure_ascii=False, indent=2)
|
||
nf.write("\n")
|
||
_os_notes.replace(temp_file, notes_file)
|
||
except Exception as e:
|
||
return web.json_response({"error": f"save notes failed: {str(e)}"}, status=500)
|
||
|
||
return web.json_response({
|
||
"success": True,
|
||
"path": notes_file,
|
||
"count": len(notes),
|
||
})
|
||
|
||
@PromptServer.instance.routes.get("/o1key/input_dir")
|
||
async def get_input_dir(request):
|
||
import os
|
||
path = os.path.abspath(folder_paths.get_input_directory())
|
||
return web.json_response({"path": path})
|
||
|
||
@PromptServer.instance.routes.get("/o1key/api_key")
|
||
async def get_api_key_route(request):
|
||
config = load_config()
|
||
key = config.get("O1KEY_API_KEY", "")
|
||
masked = ""
|
||
if key:
|
||
if len(key) > 8:
|
||
masked = key[:3] + "****" + key[-4:]
|
||
else:
|
||
masked = "****"
|
||
return web.json_response({
|
||
"has_key": bool(key),
|
||
"masked": masked,
|
||
"network_route": get_network_route(),
|
||
"network_route_options": NETWORK_ROUTE_OPTIONS,
|
||
})
|
||
|
||
@PromptServer.instance.routes.post("/o1key/config")
|
||
async def set_o1key_config(request):
|
||
data = await request.json()
|
||
route = str(data.get("network_route", "")).strip()
|
||
if route not in NETWORK_ROUTES:
|
||
return web.json_response({"error": "网络线路无效"}, status=400)
|
||
|
||
updates = {NETWORK_ROUTE_CONFIG_KEY: route}
|
||
if data.get("api_key") is not None:
|
||
new_key = str(data.get("api_key", "")).strip()
|
||
if not new_key:
|
||
return web.json_response({"error": "API Key 不能为空"}, status=400)
|
||
if "\n" in new_key or "\r" in new_key:
|
||
return web.json_response({"error": "API Key 格式无效"}, status=400)
|
||
updates["O1KEY_API_KEY"] = new_key
|
||
|
||
config = update_config(updates=updates)
|
||
key = config.get("O1KEY_API_KEY", "")
|
||
masked = key[:3] + "****" + key[-4:] if len(key) > 8 else ("****" if key else "")
|
||
return web.json_response({
|
||
"success": True,
|
||
"has_key": bool(key),
|
||
"masked": masked,
|
||
"network_route": route,
|
||
})
|
||
|
||
@PromptServer.instance.routes.post("/o1key/network_route")
|
||
async def set_network_route(request):
|
||
data = await request.json()
|
||
route = str(data.get("network_route", "")).strip()
|
||
if route not in NETWORK_ROUTES:
|
||
return web.json_response({"error": "网络线路无效"}, status=400)
|
||
update_config(updates={NETWORK_ROUTE_CONFIG_KEY: route})
|
||
return web.json_response({"success": True, "network_route": route})
|
||
|
||
@PromptServer.instance.routes.post("/o1key/api_key")
|
||
async def set_api_key_route(request):
|
||
data = await request.json()
|
||
new_key = data.get("api_key", "").strip()
|
||
if not new_key:
|
||
return web.json_response({"error": "API Key 不能为空"}, status=400)
|
||
if "\n" in new_key or "\r" in new_key:
|
||
return web.json_response({"error": "API Key 格式无效"}, status=400)
|
||
update_config(updates={"O1KEY_API_KEY": new_key})
|
||
return web.json_response({"success": True})
|
||
|
||
@PromptServer.instance.routes.post("/o1key/test_key")
|
||
async def test_api_key_route(request):
|
||
import aiohttp as _aiohttp
|
||
data = await request.json()
|
||
test_key = str(data.get("api_key") or "").strip() or get_api_key()
|
||
if not test_key:
|
||
return web.json_response({"valid": False, "error": "请先输入或保存 API Key"})
|
||
requested_route = str(data.get("network_route", "")).strip()
|
||
route = requested_route if requested_route in NETWORK_ROUTES else get_network_route()
|
||
base_url = NETWORK_ROUTES.get(route, NETWORK_ROUTES[DEFAULT_NETWORK_ROUTE])
|
||
url = f"{base_url}/v1/models"
|
||
headers = {"Authorization": f"Bearer {test_key}"}
|
||
try:
|
||
async with _aiohttp.ClientSession() as session:
|
||
async with session.get(url, headers=headers, timeout=_aiohttp.ClientTimeout(total=10)) as resp:
|
||
if resp.status == 200:
|
||
return web.json_response({"valid": True})
|
||
elif resp.status == 401:
|
||
return web.json_response({"valid": False, "error": "密钥无效或已过期"})
|
||
else:
|
||
text = await resp.text()
|
||
return web.json_response({"valid": False, "error": f"验证失败 ({resp.status})"})
|
||
except Exception as e:
|
||
return web.json_response({"valid": False, "error": f"网络错误: {str(e)}"})
|
||
|
||
@PromptServer.instance.routes.delete("/o1key/api_key")
|
||
async def delete_api_key_route(request):
|
||
update_config(remove=["O1KEY_API_KEY"])
|
||
return web.json_response({"success": True})
|
||
|
||
# === 主体(Element)代理:转发到 {base}/kling/v1/general/*,后端注入令牌 ===
|
||
_ELEMENT_PREFIX = "/kling/v1/general"
|
||
|
||
def _element_base_url(_route=None):
|
||
return NETWORK_ROUTES[get_network_route()].rstrip("/")
|
||
|
||
def _element_headers():
|
||
config = load_config()
|
||
key = config.get("O1KEY_API_KEY", "")
|
||
if not key:
|
||
return None
|
||
return {"Authorization": f"Bearer {key}"}
|
||
|
||
@PromptServer.instance.routes.get("/o1key/element/mine")
|
||
async def o1key_element_mine(request):
|
||
"""列表接口:GET /kling/v1/general/advanced-custom-elements"""
|
||
import aiohttp as _aiohttp
|
||
headers = _element_headers()
|
||
if not headers:
|
||
return web.json_response({"success": False, "message": "未配置 API Key"}, status=401)
|
||
base = _element_base_url(request.query.get("route"))
|
||
url = f"{base}{_ELEMENT_PREFIX}/advanced-custom-elements"
|
||
params = {}
|
||
# 支持分页参数
|
||
page_num = request.query.get("pageNum", "1")
|
||
page_size = request.query.get("pageSize", "100")
|
||
params["pageNum"] = page_num
|
||
params["pageSize"] = page_size
|
||
try:
|
||
async with _aiohttp.ClientSession() as session:
|
||
async with session.get(url, headers=headers, params=params,
|
||
timeout=_aiohttp.ClientTimeout(total=30)) as up:
|
||
result = await up.json()
|
||
# 新API返回: {"success": true, "data": {"code": 0, "data": [...], "total": N}, "message": ""}
|
||
# 转换为前端期望的格式: {"success": true, "data": [...]}
|
||
if result.get("success") and isinstance(result.get("data"), dict):
|
||
elements = result["data"].get("data", [])
|
||
return web.json_response({"success": True, "data": elements, "message": ""})
|
||
return web.json_response(result, status=up.status)
|
||
except Exception as e:
|
||
return web.json_response({"success": False, "message": str(e)}, status=502)
|
||
|
||
@PromptServer.instance.routes.post("/o1key/element/upload")
|
||
async def o1key_element_upload(request):
|
||
"""转发 multipart 文件上传:POST /kling/v1/general/upload
|
||
视频可达 200MB,固定总超时会截断大文件上传,改用:不限总时长 +
|
||
读空闲 120s 超时(连接卡死才超时,慢速大文件不会被一刀切断)。"""
|
||
import aiohttp as _aiohttp
|
||
headers = _element_headers()
|
||
if not headers:
|
||
return web.json_response({"success": False, "message": "未配置 API Key"}, status=401)
|
||
base = _element_base_url(request.query.get("route"))
|
||
url = f"{base}{_ELEMENT_PREFIX}/upload"
|
||
try:
|
||
reader = await request.multipart()
|
||
field = await reader.next()
|
||
if field is None or field.name != "file":
|
||
return web.json_response({"success": False, "message": "缺少 file 字段"}, status=400)
|
||
file_bytes = await field.read(decode=False)
|
||
filename = field.filename or "image.png"
|
||
form = _aiohttp.FormData()
|
||
form.add_field("file", file_bytes, filename=filename,
|
||
content_type=field.headers.get("Content-Type", "application/octet-stream"))
|
||
timeout = _aiohttp.ClientTimeout(total=None, sock_connect=30, sock_read=120)
|
||
async with _aiohttp.ClientSession() as session:
|
||
async with session.post(url, headers=headers, data=form,
|
||
timeout=timeout) as up:
|
||
data = await up.json()
|
||
return web.json_response(data, status=up.status)
|
||
except Exception as e:
|
||
return web.json_response({"success": False, "message": str(e)}, status=502)
|
||
|
||
@PromptServer.instance.routes.post("/o1key/element/create")
|
||
async def o1key_element_create(request):
|
||
"""创建主体:POST /kling/v1/general/advanced-custom-elements
|
||
|
||
新API字段映射:
|
||
- name -> element_name
|
||
- description -> element_description
|
||
- reference_type -> reference_type (image_refer / video_refer)
|
||
- frontal_image -> frontal_image
|
||
- refer_images -> refer_images
|
||
- video_list -> video_list
|
||
- element_voice_id, tag_ids, channel_id 保持不变
|
||
"""
|
||
import aiohttp as _aiohttp
|
||
import json as _json_element
|
||
headers = _element_headers()
|
||
if not headers:
|
||
return web.json_response({"success": False, "message": "未配置 API Key"}, status=401)
|
||
try:
|
||
payload = await request.json()
|
||
except Exception as e:
|
||
return web.json_response({"success": False, "message": f"请求体无效: {e}"}, status=400)
|
||
route = payload.pop("route", None)
|
||
base = _element_base_url(route)
|
||
url = f"{base}{_ELEMENT_PREFIX}/advanced-custom-elements"
|
||
|
||
# 字段映射:前端使用旧字段名,转换为新API字段名
|
||
api_payload = {}
|
||
if "name" in payload:
|
||
api_payload["element_name"] = payload["name"]
|
||
if "description" in payload:
|
||
api_payload["element_description"] = payload["description"]
|
||
# 其他字段直接透传
|
||
for key in ["reference_type", "frontal_image", "refer_images", "video_list",
|
||
"element_voice_id", "tag_ids", "channel_id"]:
|
||
if key in payload:
|
||
api_payload[key] = payload[key]
|
||
|
||
send_headers = {**headers, "Content-Type": "application/json"}
|
||
# 打印创建主体的请求信息
|
||
try:
|
||
print(f"[主体创建] 请求 URL: {url}")
|
||
print("[主体创建] 请求体: " + _json_element.dumps(api_payload, ensure_ascii=False, indent=2))
|
||
except Exception:
|
||
pass
|
||
try:
|
||
async with _aiohttp.ClientSession() as session:
|
||
async with session.post(url, headers=send_headers, json=api_payload,
|
||
timeout=_aiohttp.ClientTimeout(total=60)) as up:
|
||
data = await up.json()
|
||
# 打印创建主体的响应信息
|
||
try:
|
||
print("[主体创建] 响应体: " + _json_element.dumps(data, ensure_ascii=False, indent=2))
|
||
except Exception:
|
||
pass
|
||
return web.json_response(data, status=up.status)
|
||
except Exception as e:
|
||
return web.json_response({"success": False, "message": str(e)}, status=502)
|
||
|
||
@PromptServer.instance.routes.post("/o1key/element/refresh")
|
||
async def o1key_element_refresh(request):
|
||
"""查询主体:GET /kling/v1/general/advanced-custom-elements/{task_id}
|
||
|
||
前端传 id(数据库主键),需要先查本地库拿到 job_id(即 task_id),再查询上游。
|
||
为了简化,这里改为前端直接传 task_id(即创建时返回的 job_id)。
|
||
"""
|
||
import aiohttp as _aiohttp
|
||
import json as _json_refresh
|
||
headers = _element_headers()
|
||
if not headers:
|
||
return web.json_response({"success": False, "message": "未配置 API Key"}, status=401)
|
||
try:
|
||
payload = await request.json()
|
||
except Exception:
|
||
payload = {}
|
||
task_id = payload.get("task_id") or payload.get("id")
|
||
if not task_id:
|
||
return web.json_response({"success": False, "message": "缺少 task_id"}, status=400)
|
||
base = _element_base_url(payload.get("route"))
|
||
url = f"{base}{_ELEMENT_PREFIX}/advanced-custom-elements/{task_id}"
|
||
try:
|
||
async with _aiohttp.ClientSession() as session:
|
||
async with session.get(url, headers=headers,
|
||
timeout=_aiohttp.ClientTimeout(total=60)) as up:
|
||
data = await up.json()
|
||
# 打印查询响应
|
||
try:
|
||
print(f"[主体查询] task_id={task_id}")
|
||
print("[主体查询] 响应体: " + _json_refresh.dumps(data, ensure_ascii=False, indent=2))
|
||
except Exception:
|
||
pass
|
||
# 新API返回嵌套结构,需要提取 task_status 和 element_id
|
||
# 响应: {"success": true, "data": {"code": 0, "data": {"task_status": "succeed", "task_result": {"elements": [...]}}}}
|
||
if data.get("success") and isinstance(data.get("data"), dict):
|
||
inner = data["data"].get("data", {})
|
||
task_status = inner.get("task_status", "")
|
||
# 转换为前端期望的格式
|
||
element = {
|
||
"id": task_id,
|
||
"job_id": task_id,
|
||
"status": task_status,
|
||
"task_status": task_status,
|
||
}
|
||
if task_status == "succeed":
|
||
elements = inner.get("task_result", {}).get("elements", [])
|
||
if elements:
|
||
first = elements[0]
|
||
element["element_id"] = str(first.get("element_id", ""))
|
||
element["name"] = first.get("element_name", "")
|
||
element["frontal_image"] = first.get("element_image_list", {}).get("frontal_image", "")
|
||
elif task_status == "failed":
|
||
element["fail_reason"] = inner.get("task_status_msg", "")
|
||
return web.json_response({"success": True, "data": {"element": element, "detail": data["data"]}})
|
||
return web.json_response(data, status=up.status)
|
||
except Exception as e:
|
||
return web.json_response({"success": False, "message": str(e)}, status=502)
|
||
|
||
@PromptServer.instance.routes.post("/o1key/element/delete")
|
||
async def o1key_element_delete(request):
|
||
"""删除主体:POST /kling/v1/general/delete-advanced-elements
|
||
|
||
请求体: {"element_id": "315320838184520"}
|
||
"""
|
||
import aiohttp as _aiohttp
|
||
headers = _element_headers()
|
||
if not headers:
|
||
return web.json_response({"success": False, "message": "未配置 API Key"}, status=401)
|
||
try:
|
||
payload = await request.json()
|
||
except Exception:
|
||
payload = {}
|
||
element_id = payload.get("element_id") or payload.get("id")
|
||
if not element_id:
|
||
return web.json_response({"success": False, "message": "缺少 element_id"}, status=400)
|
||
base = _element_base_url(payload.get("route"))
|
||
url = f"{base}{_ELEMENT_PREFIX}/delete-advanced-elements"
|
||
delete_payload = {"element_id": str(element_id)}
|
||
send_headers = {**headers, "Content-Type": "application/json"}
|
||
try:
|
||
async with _aiohttp.ClientSession() as session:
|
||
async with session.post(url, headers=send_headers, json=delete_payload,
|
||
timeout=_aiohttp.ClientTimeout(total=60)) as up:
|
||
data = await up.json()
|
||
return web.json_response(data, status=up.status)
|
||
except Exception as e:
|
||
return web.json_response({"success": False, "message": str(e)}, status=502)
|
||
|
||
@PromptServer.instance.routes.get("/o1key/element/image")
|
||
async def o1key_element_image(request):
|
||
"""图片同源代理:ComfyUI 的 CSP 限制 img-src 'self',外链缩略图无法直接显示。
|
||
前端把缩略图 src 指向本路由,后端取回字节再吐给浏览器,对浏览器即同源。
|
||
仅允许 o1key 资源域,避免被当成任意 URL 抓取的 SSRF 跳板。"""
|
||
import aiohttp as _aiohttp
|
||
from urllib.parse import urlparse, unquote
|
||
raw = request.query.get("url", "")
|
||
if not raw:
|
||
return web.json_response({"success": False, "message": "缺少 url"}, status=400)
|
||
target = unquote(raw)
|
||
try:
|
||
parsed = urlparse(target)
|
||
except Exception:
|
||
parsed = None
|
||
if not parsed or parsed.scheme not in ("http", "https"):
|
||
return web.json_response({"success": False, "message": "非法 url"}, status=400)
|
||
host = (parsed.hostname or "").lower()
|
||
if not (host.endswith(".o1key.com") or host.endswith(".o1key.cn")
|
||
or host in ("o1key.com", "o1key.cn")):
|
||
return web.json_response({"success": False, "message": "不允许的图片来源"}, status=403)
|
||
try:
|
||
async with _aiohttp.ClientSession() as session:
|
||
async with session.get(target, timeout=_aiohttp.ClientTimeout(total=30)) as up:
|
||
if up.status != 200:
|
||
return web.Response(status=up.status)
|
||
body = await up.read()
|
||
ctype = up.headers.get("Content-Type", "image/jpeg").split(";")[0].strip()
|
||
return web.Response(body=body, content_type=ctype or "image/jpeg",
|
||
headers={"Cache-Control": "max-age=3600"})
|
||
except Exception as e:
|
||
return web.json_response({"success": False, "message": str(e)}, status=502)
|
||
|
||
@PromptServer.instance.routes.get("/o1key/output_history")
|
||
async def get_output_history(request):
|
||
"""读取 output 目录文件,按执行分组返回 /api/jobs 兼容格式"""
|
||
import os, json as _json
|
||
limit = int(request.query.get("limit", "200"))
|
||
offset = int(request.query.get("offset", "0"))
|
||
output_dir = os.path.abspath(folder_paths.get_output_directory())
|
||
meta_file = _get_o1key_history_meta_file(output_dir)
|
||
meta = {}
|
||
if os.path.isfile(meta_file):
|
||
try:
|
||
with open(meta_file, "r", encoding="utf-8") as mf:
|
||
meta = _json.load(mf)
|
||
except Exception:
|
||
pass
|
||
supported_ext = {'.png', '.jpg', '.jpeg', '.webp', '.gif', '.mp4', '.webm'}
|
||
# 收集所有文件并按 workflow_id 分组
|
||
all_files = []
|
||
for fname in meta.keys():
|
||
ext = os.path.splitext(fname)[1].lower()
|
||
if ext not in supported_ext:
|
||
continue
|
||
fpath = os.path.join(output_dir, fname)
|
||
if not os.path.isfile(fpath):
|
||
continue
|
||
mtime = os.path.getmtime(fpath)
|
||
media = "images" if ext in {'.png','.jpg','.jpeg','.webp','.gif'} else "video"
|
||
all_files.append({"name": fname, "mtime": mtime, "media": media})
|
||
# 按 workflow_id 分组(同一次执行合并为一个 job)
|
||
groups = {}
|
||
for f in all_files:
|
||
m = meta.get(f["name"], {})
|
||
wid = m.get("workflow_id")
|
||
if wid:
|
||
groups.setdefault(wid, []).append((f, m))
|
||
# 构建 job 列表
|
||
jobs = []
|
||
for wid, items in groups.items():
|
||
items.sort(key=lambda x: x[0]["mtime"], reverse=True)
|
||
latest = items[0]
|
||
f, m = latest
|
||
start_ms = int(m.get("start_time", f["mtime"]) * 1000)
|
||
end_ms = int(m.get("end_time", f["mtime"]) * 1000)
|
||
jobs.append({
|
||
"id": wid,
|
||
"status": "completed",
|
||
"create_time": start_ms,
|
||
"execution_start_time": start_ms,
|
||
"execution_end_time": end_ms,
|
||
"preview_output": {
|
||
"filename": f["name"],
|
||
"subfolder": "",
|
||
"type": "output",
|
||
"nodeId": "0",
|
||
"mediaType": f["media"],
|
||
},
|
||
"outputs_count": len(items),
|
||
"execution_error": None,
|
||
"workflow_id": wid,
|
||
})
|
||
# 按时间倒序排列,分页
|
||
jobs.sort(key=lambda x: x["create_time"], reverse=True)
|
||
total = len(jobs)
|
||
page = jobs[offset:offset+limit]
|
||
return web.json_response({
|
||
"jobs": page,
|
||
"pagination": {"offset": offset, "limit": limit, "total": total, "has_more": offset + limit < total}
|
||
})
|
||
|
||
@PromptServer.instance.routes.get("/o1key/output_workflow")
|
||
async def get_output_workflow(request):
|
||
"""从 PNG 元数据中读取工作流,供前端恢复使用"""
|
||
import os, struct, json as _json
|
||
filename = request.query.get("filename", "")
|
||
if not filename:
|
||
return web.json_response({"error": "missing filename"}, status=400)
|
||
output_dir = os.path.abspath(folder_paths.get_output_directory())
|
||
fpath = os.path.join(output_dir, filename)
|
||
if not os.path.isfile(fpath) or not fpath.lower().endswith(".png"):
|
||
return web.json_response({"error": "file not found"}, status=404)
|
||
workflow = None
|
||
prompt_data = None
|
||
try:
|
||
with open(fpath, "rb") as pf:
|
||
pf.read(8) # PNG signature
|
||
while True:
|
||
raw = pf.read(8)
|
||
if len(raw) < 8:
|
||
break
|
||
length = struct.unpack(">I", raw[:4])[0]
|
||
chunk_type = raw[4:8]
|
||
data = pf.read(length)
|
||
pf.read(4) # CRC
|
||
if chunk_type == b"tEXt":
|
||
key, val = data.split(b"\x00", 1)
|
||
k = key.decode("ascii", errors="replace")
|
||
if k == "workflow":
|
||
workflow = _json.loads(val)
|
||
elif k == "prompt":
|
||
prompt_data = _json.loads(val)
|
||
elif chunk_type == b"IEND":
|
||
break
|
||
except Exception:
|
||
pass
|
||
return web.json_response({"workflow": workflow, "prompt": prompt_data})
|
||
|
||
@PromptServer.instance.routes.get("/o1key/job_detail/{job_id}")
|
||
async def get_job_detail(request):
|
||
"""根据 job_id 返回当前端口持久化历史中的 job 详情"""
|
||
import os, struct, json as _json
|
||
job_id = request.match_info["job_id"]
|
||
output_dir = os.path.abspath(folder_paths.get_output_directory())
|
||
meta_file = _get_o1key_history_meta_file(output_dir)
|
||
meta = {}
|
||
if os.path.isfile(meta_file):
|
||
try:
|
||
with open(meta_file, "r", encoding="utf-8") as mf:
|
||
meta = _json.load(mf)
|
||
except Exception:
|
||
pass
|
||
supported_ext = {'.png', '.jpg', '.jpeg', '.webp', '.gif', '.mp4', '.webm'}
|
||
# 只在当前端口的持久化记录中查找该 job 的文件
|
||
matched_files = []
|
||
for fname in meta.keys():
|
||
ext = os.path.splitext(fname)[1].lower()
|
||
if ext not in supported_ext:
|
||
continue
|
||
m = meta.get(fname, {})
|
||
if m.get("workflow_id") == job_id:
|
||
matched_files.append(fname)
|
||
if not matched_files:
|
||
return web.json_response({"error": "not found"}, status=404)
|
||
# 用最新文件作为代表
|
||
matched_files.sort(key=lambda f: os.path.getmtime(os.path.join(output_dir, f)), reverse=True)
|
||
target_file = matched_files[0]
|
||
fpath = os.path.join(output_dir, target_file)
|
||
m = meta.get(target_file, {})
|
||
mtime = os.path.getmtime(fpath)
|
||
start_ms = int(m.get("start_time", mtime) * 1000)
|
||
end_ms = int(m.get("end_time", mtime) * 1000)
|
||
ext = os.path.splitext(target_file)[1].lower()
|
||
media = "images" if ext in {'.png','.jpg','.jpeg','.webp','.gif'} else "video"
|
||
# 读取 PNG 工作流元数据
|
||
workflow = None
|
||
if ext == ".png":
|
||
try:
|
||
with open(fpath, "rb") as pf:
|
||
pf.read(8)
|
||
while True:
|
||
raw = pf.read(8)
|
||
if len(raw) < 8:
|
||
break
|
||
length = struct.unpack(">I", raw[:4])[0]
|
||
chunk_type = raw[4:8]
|
||
data = pf.read(length)
|
||
pf.read(4)
|
||
if chunk_type == b"tEXt":
|
||
key, val = data.split(b"\x00", 1)
|
||
k = key.decode("ascii", errors="replace")
|
||
if k == "workflow":
|
||
workflow = _json.loads(val)
|
||
elif chunk_type == b"IEND":
|
||
break
|
||
except Exception:
|
||
pass
|
||
# 构建 outputs:包含该执行的所有文件
|
||
outputs = {}
|
||
for i, fname in enumerate(matched_files):
|
||
e = os.path.splitext(fname)[1].lower()
|
||
mt = "images" if e in {'.png','.jpg','.jpeg','.webp','.gif'} else "gifs"
|
||
outputs.setdefault(str(i), {}).setdefault(mt, []).append(
|
||
{"filename": fname, "subfolder": "", "type": "output"}
|
||
)
|
||
job_detail = {
|
||
"id": job_id,
|
||
"status": "completed",
|
||
"create_time": start_ms,
|
||
"execution_start_time": start_ms,
|
||
"execution_end_time": end_ms,
|
||
"preview_output": {
|
||
"filename": target_file,
|
||
"subfolder": "",
|
||
"type": "output",
|
||
"nodeId": "0",
|
||
"mediaType": media,
|
||
},
|
||
"outputs_count": len(matched_files),
|
||
"execution_error": None,
|
||
"workflow_id": job_id,
|
||
"workflow": {
|
||
"extra_data": {
|
||
"extra_pnginfo": {"workflow": workflow}
|
||
}
|
||
} if workflow else None,
|
||
"outputs": outputs,
|
||
}
|
||
return web.json_response(job_detail)
|
||
|
||
@PromptServer.instance.routes.post("/o1key/delete_history")
|
||
async def delete_history_item(request):
|
||
"""删除持久化历史记录及对应的输出文件"""
|
||
import os, json as _json
|
||
body = await request.json()
|
||
job_ids = body.get("delete", [])
|
||
if not job_ids:
|
||
return web.json_response({"success": False, "error": "missing ids"}, status=400)
|
||
output_dir = os.path.abspath(folder_paths.get_output_directory())
|
||
meta_file = _get_o1key_history_meta_file(output_dir)
|
||
meta = {}
|
||
if os.path.isfile(meta_file):
|
||
try:
|
||
with open(meta_file, "r", encoding="utf-8") as mf:
|
||
meta = _json.load(mf)
|
||
except Exception:
|
||
pass
|
||
deleted_files = []
|
||
for job_id in job_ids:
|
||
files_to_remove = []
|
||
for fname, m in list(meta.items()):
|
||
if m.get("workflow_id") == job_id:
|
||
files_to_remove.append(fname)
|
||
for fname in files_to_remove:
|
||
meta.pop(fname, None)
|
||
fpath = os.path.join(output_dir, fname)
|
||
if os.path.isfile(fpath):
|
||
try:
|
||
os.remove(fpath)
|
||
deleted_files.append(fname)
|
||
except Exception:
|
||
pass
|
||
try:
|
||
with open(meta_file, "w", encoding="utf-8") as mf:
|
||
_json.dump(meta, mf, ensure_ascii=False)
|
||
except Exception:
|
||
pass
|
||
return web.json_response({"success": True, "deleted": deleted_files})
|
||
|
||
# === 图片生成提示词优化(服务端读取参考图,避免前端接触 API Key) ===
|
||
@PromptServer.instance.routes.post("/o1key/image/prompt-optimize")
|
||
async def optimize_o1key_image_prompt(request):
|
||
import aiohttp as _aiohttp
|
||
|
||
try:
|
||
data = await request.json()
|
||
if not isinstance(data, dict):
|
||
raise ValueError("请求体必须是对象")
|
||
prompt = data.get("prompt", "")
|
||
references = data.get("references", [])
|
||
api_key = get_api_key() or ""
|
||
if not api_key:
|
||
return web.json_response({"error": "未配置 API Key"}, status=401)
|
||
|
||
timeout = _aiohttp.ClientTimeout(total=PROMPT_OPTIMIZER_TIMEOUT_SECONDS + 15)
|
||
async with _aiohttp.ClientSession(timeout=timeout) as session:
|
||
optimized = await optimize_image_prompt(
|
||
session,
|
||
NETWORK_ROUTES[get_network_route()],
|
||
api_key,
|
||
prompt,
|
||
references,
|
||
folder_paths.get_input_directory(),
|
||
)
|
||
return web.json_response(
|
||
{
|
||
"prompt": optimized,
|
||
"model": "gpt-5.6-sol",
|
||
"reasoning_effort": "high",
|
||
},
|
||
headers={"Cache-Control": "no-store"},
|
||
)
|
||
except ValueError as exc:
|
||
return web.json_response({"error": str(exc)}, status=400)
|
||
except RuntimeError as exc:
|
||
return web.json_response({"error": str(exc)}, status=502)
|
||
except Exception:
|
||
return web.json_response({"error": "提示词优化失败,请稍后重试"}, status=500)
|
||
|
||
# === 视频生成 AI帮写(独立视频预设,仅分析安全的 input 图片描述) ===
|
||
@PromptServer.instance.routes.post("/o1key/video/prompt-write")
|
||
async def write_o1key_video_prompt(request):
|
||
import aiohttp as _aiohttp
|
||
|
||
try:
|
||
data = await request.json()
|
||
if not isinstance(data, dict):
|
||
raise ValueError("请求体必须是对象")
|
||
api_key = get_api_key() or ""
|
||
if not api_key:
|
||
return web.json_response({"error": "未配置 API Key"}, status=401)
|
||
|
||
context = {
|
||
"generation_mode": data.get("generation_mode", "text"),
|
||
"duration": data.get("duration", "auto"),
|
||
"aspect_ratio": data.get("aspect_ratio", "auto"),
|
||
"generate_audio": data.get("generate_audio", False),
|
||
"reference_video_count": data.get("reference_video_count", 0),
|
||
"reference_audio_count": data.get("reference_audio_count", 0),
|
||
}
|
||
timeout = _aiohttp.ClientTimeout(total=PROMPT_OPTIMIZER_TIMEOUT_SECONDS + 15)
|
||
async with _aiohttp.ClientSession(timeout=timeout) as session:
|
||
written = await write_video_prompt(
|
||
session,
|
||
NETWORK_ROUTES[get_network_route()],
|
||
api_key,
|
||
data.get("prompt", ""),
|
||
data.get("references", []),
|
||
folder_paths.get_input_directory(),
|
||
context,
|
||
)
|
||
return web.json_response(
|
||
{
|
||
"prompt": written,
|
||
"model": "gpt-5.6-sol",
|
||
"reasoning_effort": "high",
|
||
"preset": "video-default",
|
||
},
|
||
headers={"Cache-Control": "no-store"},
|
||
)
|
||
except ValueError as exc:
|
||
return web.json_response({"error": str(exc)}, status=400)
|
||
except RuntimeError as exc:
|
||
return web.json_response({"error": str(exc)}, status=502)
|
||
except Exception:
|
||
return web.json_response({"error": "视频 AI帮写失败,请稍后重试"}, status=500)
|
||
|
||
# === AI 聊天代理(流式 SSE 透传) ===
|
||
@PromptServer.instance.routes.post("/o1key/chat/completions")
|
||
async def chat_completions_proxy(request):
|
||
import aiohttp as _aiohttp
|
||
import json as _cjson
|
||
|
||
data = await request.json()
|
||
config = load_config()
|
||
api_key = config.get("O1KEY_API_KEY", "")
|
||
if not api_key:
|
||
return web.json_response({"error": "未配置 API Key"}, status=401)
|
||
|
||
base_url = NETWORK_ROUTES[get_network_route()]
|
||
model = data.get("model", "gpt-6-sol")
|
||
messages = data.get("messages", [])
|
||
if not isinstance(messages, list) or not messages:
|
||
return web.json_response({"error": "缺少对话内容"}, status=400)
|
||
|
||
try:
|
||
messages = expand_xlsx_attachments(messages)
|
||
except ValueError as exc:
|
||
return web.json_response({"error": str(exc)}, status=400)
|
||
|
||
url = f"{base_url}/v1/chat/completions"
|
||
headers = {
|
||
"Content-Type": "application/json",
|
||
"Authorization": f"Bearer {api_key}",
|
||
}
|
||
reasoning = data.get("reasoning_effort", "high")
|
||
if reasoning not in ("low", "medium", "high"):
|
||
reasoning = "high"
|
||
body = {"model": model, "messages": messages, "stream": True}
|
||
if model == "claude-fable-5":
|
||
budgets = {"low": 2048, "medium": 8192, "high": 16384}
|
||
budget = budgets[reasoning]
|
||
body["thinking"] = {"type": "enabled", "budget_tokens": budget}
|
||
body["max_tokens"] = budget + 8192
|
||
elif model in ("gpt-5.5", "gpt-5.6-sol", "gpt-6-astra", "gpt-6-sol", "gemini-3.1-pro-preview"):
|
||
body["reasoning_effort"] = reasoning
|
||
|
||
search_trace = None
|
||
timeout = _aiohttp.ClientTimeout(total=120)
|
||
async with _aiohttp.ClientSession(timeout=timeout) as session:
|
||
if data.get("web_search") is True:
|
||
raw_query = extract_search_query(messages)
|
||
if raw_query:
|
||
query = await rewrite_search_query(session, base_url, api_key, raw_query) or raw_query
|
||
try:
|
||
results = await web_search(session, query)
|
||
search_trace = {
|
||
"query": query,
|
||
"results": [
|
||
{"title": item["title"], "url": item["url"]}
|
||
for item in results
|
||
],
|
||
}
|
||
messages = list(messages)
|
||
messages.insert(max(0, len(messages) - 1), {
|
||
"role": "system",
|
||
"content": build_search_context(query, results),
|
||
})
|
||
body["messages"] = messages
|
||
except Exception as exc:
|
||
search_trace = {"query": query, "results": [], "error": str(exc)}
|
||
|
||
resp = web.StreamResponse(
|
||
status=200, reason="OK",
|
||
headers={
|
||
"Content-Type": "text/event-stream",
|
||
"Cache-Control": "no-cache",
|
||
"X-Accel-Buffering": "no",
|
||
}
|
||
)
|
||
await resp.prepare(request)
|
||
|
||
try:
|
||
if search_trace:
|
||
event = _cjson.dumps({"o1key_search": search_trace}, ensure_ascii=False)
|
||
await resp.write(f"data: {event}\n\n".encode("utf-8"))
|
||
async with session.post(url, headers=headers, json=body) as upstream:
|
||
if upstream.status != 200:
|
||
err = await upstream.text()
|
||
event = _cjson.dumps({"error": err}, ensure_ascii=False)
|
||
await resp.write(f"data: {event}\n\n".encode("utf-8"))
|
||
await resp.write(b"data: [DONE]\n\n")
|
||
return resp
|
||
async for chunk in upstream.content.iter_any():
|
||
await resp.write(chunk)
|
||
except Exception as e:
|
||
event = _cjson.dumps({"error": str(e)}, ensure_ascii=False)
|
||
await resp.write(f"data: {event}\n\n".encode("utf-8"))
|
||
await resp.write(b"data: [DONE]\n\n")
|
||
|
||
return resp
|
||
# === 执行事件 Hook:持久化耗时元数据 ===
|
||
import time as _time, json as _json2, os as _os
|
||
_execution_tracker = {}
|
||
|
||
_orig_send_sync = PromptServer.instance.send_sync
|
||
|
||
def _patched_send_sync(event, data, *args, **kwargs):
|
||
try:
|
||
if event == "execution_start":
|
||
pid = data.get("prompt_id", "")
|
||
if pid:
|
||
_execution_tracker[pid] = {"start": _time.time(), "outputs": []}
|
||
elif event == "executed":
|
||
pid = data.get("prompt_id", "")
|
||
output = data.get("output") or {}
|
||
if pid and pid in _execution_tracker:
|
||
for img in output.get("images", []) + output.get("gifs", []):
|
||
if img.get("type") == "output" and img.get("filename"):
|
||
_execution_tracker[pid]["outputs"].append(img["filename"])
|
||
elif event == "executing" and data.get("node") is None:
|
||
pid = data.get("prompt_id", "")
|
||
tracker = _execution_tracker.pop(pid, None)
|
||
if tracker and tracker["outputs"]:
|
||
end_time = _time.time()
|
||
start_time = tracker["start"]
|
||
output_dir = _os.path.abspath(folder_paths.get_output_directory())
|
||
meta_file = _get_o1key_history_meta_file(output_dir)
|
||
meta = {}
|
||
if _os.path.isfile(meta_file):
|
||
try:
|
||
with open(meta_file, "r", encoding="utf-8") as mf:
|
||
meta = _json2.load(mf)
|
||
except Exception:
|
||
pass
|
||
for fname in tracker["outputs"]:
|
||
meta[fname] = {
|
||
"start_time": start_time,
|
||
"end_time": end_time,
|
||
"outputs_count": len(tracker["outputs"]),
|
||
"workflow_id": pid,
|
||
}
|
||
try:
|
||
with open(meta_file, "w", encoding="utf-8") as mf:
|
||
_json2.dump(meta, mf, ensure_ascii=False)
|
||
except Exception:
|
||
pass
|
||
except Exception:
|
||
pass
|
||
return _orig_send_sync(event, data, *args, **kwargs)
|
||
|
||
PromptServer.instance.send_sync = _patched_send_sync
|
||
|
||
except Exception:
|
||
pass
|