From f115d5289732b646c74f8bd75bf3f2789e82a503 Mon Sep 17 00:00:00 2001 From: bedovyy Date: Sat, 13 Dec 2025 00:48:17 +0900 Subject: [PATCH] initial release --- __init__.py | 24 +++++++++ env.py | 15 ++++++ nodes.py | 117 ++++++++++++++++++++++++++++++++++++++++++++ requirements.txt | 1 + routes.py | 28 +++++++++++ web/js/getmodels.js | 38 ++++++++++++++ 6 files changed, 223 insertions(+) create mode 100644 __init__.py create mode 100644 env.py create mode 100644 nodes.py create mode 100644 requirements.txt create mode 100644 routes.py create mode 100644 web/js/getmodels.js diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..1137f74 --- /dev/null +++ b/__init__.py @@ -0,0 +1,24 @@ +from typing_extensions import override +from comfy_api.latest import ComfyExtension, io +from .nodes import * +from .routes import * + +from dotenv import load_dotenv +import folder_paths + +env_path = os.path.join(folder_paths.base_path, ".env") +if os.path.exists(env_path): + load_dotenv(env_path) + +class LLMHelperExtension(ComfyExtension): + @override + async def get_node_list(self) -> list[type[io.ComfyNode]]: + return [ + GetModels, + PostModelsUnload, + ] + +async def comfy_entrypoint() -> LLMHelperExtension: + return LLMHelperExtension() + +WEB_DIRECTORY = "./web" diff --git a/env.py b/env.py new file mode 100644 index 0000000..093982a --- /dev/null +++ b/env.py @@ -0,0 +1,15 @@ +import os +from dotenv import dotenv_values +import folder_paths + +_ENV_PATH = os.path.join(folder_paths.base_path, ".env") +_ENV = dotenv_values(_ENV_PATH) + +def get_env_keys(): + return list(_ENV.keys()) + +def get_env(key: str, default=None): + return _ENV.get(key, default) + +def get_envs(): + return dict(_ENV) diff --git a/nodes.py b/nodes.py new file mode 100644 index 0000000..f5b2ede --- /dev/null +++ b/nodes.py @@ -0,0 +1,117 @@ +import os +import requests +import folder_paths +from comfy_api.latest import io + +from .env import get_env_keys, get_env + +class GetModels(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + env_vars = get_env_keys() + env_vars.insert(0, "/* no api key */") + return io.Schema( + is_output_node=True, + node_id="LLMHelper_GetModels", + display_name="LLMHelper GET /models", + category="LLMHelper", + description="Get models.", + inputs=[ + io.String.Input( + id="base_url", + display_name="Base URL", + tooltip="The base URL to use for /models/unload", + placeholder="http(s)://host[:port]", + default="http://localhost:8000", + ), + io.Combo.Input( + id="env_var", + display_name=".env API key", + tooltip="The environment variable for API key to use.", + options=env_vars, + ), + io.Combo.Input( + id="model_name", + display_name="Model name", + tooltip="Select model.", + options=["set url and click update"] + ), + ], + outputs=[ + io.String.Output(id="output_base_url", display_name="BASE_URL"), + io.String.Output(id="output_api_key", display_name="API_KEY"), + io.String.Output(id="output_model_name", display_name="MODEL_NAME"), + ] + ) + + @classmethod + def validate_inputs(cls, base_url) -> bool | str: + if base_url == "": + return "base_url must be specified" + return True + + @classmethod + def execute(cls, base_url, env_var, model_name) -> io.NodeOutput: + api_key = get_env(env_var, "") + return io.NodeOutput(base_url.rstrip("/"), api_key, model_name) + +class PostModelsUnload(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id="LLMHelper_PostModelsUnload", + display_name="LLMHelper POST /models/unload", + category="LLMHelper", + description="Unload model.", + inputs=[ + io.AnyType.Input( + id="input_any", + display_name="*", + tooltip="connect any to run the node" + ), + io.String.Input( + id="base_url", + display_name="Base URL", + tooltip="The base URL to use for /models/unload", + placeholder="http(s)://host[:port]", + default="http://localhost:8000", + ), + io.String.Input( + id="api_key", + display_name="API Key", + tooltip="The API key to use.", + ), + io.String.Input( + id="model_name", + display_name="Model name", + tooltip="The model nae to unload. leave empty if you use it for llama-swap", + ), + ], + outputs=[ + io.AnyType.Output( + id="output_any", + tooltip="connect any to bypass", + ), + ], + ) + @classmethod + def validate_inputs(cls, base_url) -> bool | str: + if base_url == "": + return "base_url must be specified" + return True + + @classmethod +# def fingerprint_inputs(cls, **kwargs) -> str: +# return str(time.time()) # force run + + @classmethod + def execute(cls, input_any, base_url, api_key, model_name) -> io.NodeOutput: + modified_base_url = base_url.rstrip("/").removesuffix("/v1") + url = f"{modified_base_url}/models/unload" + headers = {} + if api_key != "": + headers["Authorization"] = f"Bearer {api_key}" + data = { "model": model_name, "model_name": model_name } + resp = requests.post(url, headers=headers, json=data, timeout=1) + return io.NodeOutput(input_any) + diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..566cccb --- /dev/null +++ b/requirements.txt @@ -0,0 +1 @@ +python-dotenv diff --git a/routes.py b/routes.py new file mode 100644 index 0000000..74a5885 --- /dev/null +++ b/routes.py @@ -0,0 +1,28 @@ +import os +import requests +from aiohttp import web +from server import PromptServer +from .env import get_env + +routes = PromptServer.instance.routes +@routes.post('/llmhelper/models') +async def post_models(request): + data = await request.json() + base_url = data["base_url"] + api_key = get_env(data["env_var"], "") + headers = {} + if api_key != "": + headers["Authorization"] = f"Bearer {api_key}" + response = { "models": ["model not found"] } + try: + resp = requests.get(f"{base_url}/models", headers=headers, timeout=1) + resp.raise_for_status() + json = resp.json() + if "data" in json: + ids = [item["id"] for item in json["data"]] + response["models"] = ids + except requests.exceptions.RequestException as e: + r = getattr(e, "response", None) + response["models"] = [f"{r.status_code}:{r.reason}"] + + return web.json_response(response) diff --git a/web/js/getmodels.js b/web/js/getmodels.js new file mode 100644 index 0000000..57a65b8 --- /dev/null +++ b/web/js/getmodels.js @@ -0,0 +1,38 @@ +const { app } = window.comfyAPI.app; +const { api } = window.comfyAPI.api; + +app.registerExtension({ + name: "LLMHelper.getmodels", + async beforeRegisterNodeDef(nodeType, nodeData, app) { + if (!nodeData?.category?.startsWith("LLMHelper")) { return; } + + if (nodeData.name == "LLMHelper_GetModels") { + nodeType.prototype.onConnectInput = function () { + app.extensionManager.toast.add({ + severity: "info", + summary: nodeData.display_name, + detail: "This node cannot have input connections.", + life: 5000, + }); + return false; + } //prevent input connection + nodeType.prototype.onNodeCreated = function () { + this.addWidget("button", "Update model names", null, async () => { + const data = { + base_url: this.widgets.find(w => w.name === "base_url")["value"], + env_var: this.widgets.find(w => w.name === "env_var")["value"], + }; + const resp = await api.fetchApi("/llmhelper/models", { method: "POST", body: JSON.stringify(data) }); + const models = (await resp.json()).models; + if (models) { + const model_name_widget = this.widgets.find(w => w.name === "model_name"); + model_name_widget["options"]["values"] = models; + if (!models.includes(model_name_widget["value"])) { + model_name_widget["value"] = models[0]; + } + } + }) + } + } + }, +})