initial release
This commit is contained in:
+24
@@ -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"
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
python-dotenv
|
||||
@@ -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)
|
||||
@@ -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];
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
},
|
||||
})
|
||||
Reference in New Issue
Block a user