diff --git a/routes.py b/routes.py index b3938ba..91d9c89 100644 --- a/routes.py +++ b/routes.py @@ -1,9 +1,11 @@ -import requests +import asyncio import aiohttp from server import PromptServer from .env import get_env +API_TIMEOUT = 5 + routes = PromptServer.instance.routes @routes.post('/api/llmhelper/models') async def post_model_list(request): @@ -18,7 +20,7 @@ async def post_model_list(request): response = { "models": ["model not found"] } try: - timeout = aiohttp.ClientTimeout(total=1) + timeout = aiohttp.ClientTimeout(total=API_TIMEOUT) async with aiohttp.ClientSession(timeout=timeout) as session: async with session.get(f"{base_url}/models", headers=headers) as resp: @@ -28,9 +30,11 @@ async def post_model_list(request): source = json_data.get("data") or json_data.get("models") or [] response["models"] = [item.get("id") or item.get("name") for item in source] + except asyncio.TimeoutError: + response["models"] = [f"Error:Request Timeout ({API_TIMEOUT}s)"] except aiohttp.ClientResponseError as e: - response["models"] = [f"{e.status}:{e.message}"] + response["models"] = [f"{e.status}:{e.message or 'API Error'}"] except Exception as e: - response["models"] = [f"Error: {str(e)}"] + response["models"] = [f"Error:{str(e)}"] return aiohttp.web.json_response(response) diff --git a/web/js/getmodels.js b/web/js/getmodels.js index e563b49..d652dbd 100644 --- a/web/js/getmodels.js +++ b/web/js/getmodels.js @@ -9,15 +9,17 @@ app.registerExtension({ async nodeCreated(node) { if (node.comfyClass !== TARGET) return; - const base_url_widget = node.widgets.find(w => w.name === "base_url"); - const env_var_widget = node.widgets.find(w => w.name === "env_var"); - const model_name_widget = node.widgets.find(w => w.name === "model_name"); + const baseUrlWidget = node.widgets.find(w => w.name === "base_url"); + const envVarWidget = node.widgets.find(w => w.name === "env_var"); + const modelNameWidget = node.widgets.find(w => w.name === "model_name"); node.addWidget("button", "Update model names", null, async () => { const data = { - base_url: base_url_widget.value, - env_var: env_var_widget.value, + base_url: baseUrlWidget.value, + env_var: envVarWidget.value, }; + const originalModelName = modelNameWidget.value + modelNameWidget.value = "Updating..." const resp = await api.fetchApi("/llmhelper/models", { method: "POST", headers: {"Content-Type": "application/json"}, @@ -25,10 +27,8 @@ app.registerExtension({ }); const models = (await resp.json()).models; if (models) { - model_name_widget.options.values = models; - if (!models.includes(model_name_widget.value)) { - model_name_widget.value = models[0]; - } + modelNameWidget.options.values = models; + modelNameWidget.value = models.includes(originalModelName) ? originalModelName : models[0]; } });