initial release

This commit is contained in:
bedovyy
2025-12-13 00:48:17 +09:00
parent 1f3b0f21c3
commit f115d52897
6 changed files with 223 additions and 0 deletions
+24
View File
@@ -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"
+15
View File
@@ -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)
+117
View File
@@ -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)
+1
View File
@@ -0,0 +1 @@
python-dotenv
+28
View File
@@ -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)
+38
View File
@@ -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];
}
}
})
}
}
},
})