Model info and res preset fix

This commit is contained in:
2026-07-26 22:42:20 -07:00
parent e9ee89f8a3
commit aa59df60d9
3 changed files with 27 additions and 8 deletions
+1 -1
View File
@@ -151,7 +151,7 @@ TOOL_REGISTRY = [
"resolution_preset": {"type": "string", "description": "A named resolution preset (e.g., 'square', 'portrait'). Available options depend on the model."},
"cfg_scale": {"type": "number", "description": "CFG scale for prompt adherence. Usually should be left omitted to select the default."},
},
"required": ["model_name", "prompt"],
"required": ["model_name", "prompt", "resolution_preset"],
},
handler=generate_image_handler
),
+14 -6
View File
@@ -1,6 +1,8 @@
import requests
import base64
import asyncio
import config
from typing import Any, Dict, List
from tools.utils import ToolError, load_toml
@@ -9,7 +11,7 @@ async def ensure_model_state(model_name: str, model_preset: Dict[str, Any]):
Checks the current server state and switches model/modules if they differ from the preset.
"""
try:
config_resp = requests.get(f"{config.SD_URL}/config", timeout=10)
config_resp = await asyncio.to_thread(requests.get, f"{config.SD_URL}/config", timeout=10)
config_resp.raise_for_status()
cfg_data = config_resp.json()
components = cfg_data.get("components", [])
@@ -28,7 +30,8 @@ async def ensure_model_state(model_name: str, model_preset: Dict[str, Any]):
if active_ckpt != target_ckpt:
preset = model_preset.get("preset", "xl")
requests.post(
await asyncio.to_thread(
requests.post,
f"{config.SD_URL}/api/predict/checkpoint_change",
json={"data": [target_ckpt, preset]},
timeout=30
@@ -40,7 +43,8 @@ async def ensure_model_state(model_name: str, model_preset: Dict[str, Any]):
if active_modules != target_modules:
preset = model_preset.get("preset", "xl")
requests.post(
await asyncio.to_thread(
requests.post,
f"{config.SD_URL}/api/predict/modules_change",
json={"data": [target_modules, preset]},
timeout=30
@@ -64,9 +68,13 @@ async def handle(args: Dict[str, Any]) -> List[Dict[str, Any]]:
if not model_name:
raise ToolError("The 'model_name' argument is required. Use get_model_info to see available models.")
res_preset_name = args.get("resolution_preset")
if not res_preset_name:
raise ToolError("The 'resolution_preset' argument is required. Use get_model_info to see available resolutions for the chosen model.")
# 2. FETCH CURRENT SERVER DEFAULTS & MAP LABELS
try:
info_resp = requests.get(f"{config.SD_URL}/info", timeout=10)
info_resp = await asyncio.to_thread(requests.get, f"{config.SD_URL}/info", timeout=10)
info_resp.raise_for_status()
info_data = info_resp.json()
params_info = info_data["named_endpoints"]["/txt2img"]["parameters"]
@@ -161,7 +169,7 @@ async def handle(args: Dict[str, Any]) -> List[Dict[str, Any]]:
# 6. EXECUTE GENERATION
try:
gen_payload = {"data": payload}
gen_resp = requests.post(f"{config.SD_URL}/api/txt2img", json=gen_payload, timeout=300)
gen_resp = await asyncio.to_thread(requests.post, f"{config.SD_URL}/api/txt2img", json=gen_payload, timeout=300)
gen_resp.raise_for_status()
res_data = gen_resp.json()
@@ -184,7 +192,7 @@ async def handle(args: Dict[str, Any]) -> List[Dict[str, Any]]:
proxy_url = f"http://localhost:{config.PORT}/file{raw_path}"
# Download the image and convert to base64 for the AI's vision
img_resp = requests.get(sd_img_url, timeout=30)
img_resp = await asyncio.to_thread(requests.get, sd_img_url, timeout=30)
img_resp.raise_for_status()
b64_data = base64.b64encode(img_resp.content).decode('utf-8')
+12 -1
View File
@@ -26,7 +26,18 @@ async def handle(args: Dict[str, Any]) -> Dict[str, Any]:
}
if model_name not in models:
raise ToolError(f"Model '{model_name}' not found in presets. Available models: {', '.join(models.keys())}")
# Return the full catalog if the specific model isn't found
catalog = []
for name, data in models.items():
catalog.append({
"name": name,
"description": data.get("description", "No description provided.")
})
return {
"text": f"Model '{model_name}' not found.\n\nAvailable models:\n\n" +
"\n".join([f"- {m['name']}: {m['description']}" for m in catalog]) +
"\n\nTo get a detailed prompting guide and available resolutions for a specific model, call this tool again with the 'model_name' argument."
}
model_data = models[model_name]
res_set_name = model_data.get("Resolution Set")