Stable Diffusion WebUI integration. File CORP proxy.

This commit is contained in:
moosecrap 2026-07-25 14:36:43 -07:00
parent 449a41742d
commit f808f4d599
6 changed files with 344 additions and 1 deletions

27
main.py
View File

@ -5,7 +5,7 @@ import signal
import os import os
from starlette.applications import Starlette from starlette.applications import Starlette
from starlette.routing import Route from starlette.routing import Route
from starlette.responses import Response, StreamingResponse from starlette.responses import Response, StreamingResponse, FileResponse
from starlette.middleware.cors import CORSMiddleware from starlette.middleware.cors import CORSMiddleware
from mcp_logic import MCPServer from mcp_logic import MCPServer
@ -75,10 +75,35 @@ async def messages_endpoint(request):
await mcp_logic.send_to_session(sid, {"jsonrpc": "2.0", "id": req_id, "result": result}) await mcp_logic.send_to_session(sid, {"jsonrpc": "2.0", "id": req_id, "result": result})
return Response(status_code=202) return Response(status_code=202)
async def file_endpoint(request):
"""
Proxy endpoint to serve files from disk with CORP/CORS headers to bypass browser restrictions.
"""
path = request.path_params.get("path")
if not path:
return Response("Path not provided", status_code=400)
# Ensure the path is absolute (SDFiles are usually in /tmp/gradio/...)
# If the path doesn't start with /, we assume it's relative to root for simplicity in this context
full_path = path if path.startswith("/") else f"/{path}"
if not os.path.exists(full_path):
return Response(f"File not found: {full_path}", status_code=404)
return FileResponse(
full_path,
headers={
"Cross-Origin-Resource-Policy": "cross-origin",
"Access-Control-Allow-Origin": "*",
"Access-Control-Allow-Methods": "GET"
}
)
app = Starlette( app = Starlette(
routes=[ routes=[
Route("/sse", endpoint=sse_endpoint), Route("/sse", endpoint=sse_endpoint),
Route("/messages", endpoint=messages_endpoint, methods=["POST"]), Route("/messages", endpoint=messages_endpoint, methods=["POST"]),
Route("/file/{path:path}", endpoint=file_endpoint),
] ]
) )

33
model_presets.toml Normal file
View File

@ -0,0 +1,33 @@
# Model Presets
# Format: ["Model Name"]
# Description and Guide are mandatory.
# Other keys should be the Labels found in the /info endpoint.
["noobaiXLNAIXL_vPred10Version"]
description = "High-quality anime model, requires tag-based prompting."
guide = """
Detailed Prompting Guide:
- Style: Tag-based (Danbooru). Use comma-separated tags rather than natural language.
- Quality: Use 'masterpiece, best quality' for high fidelity.
- Negatives: Use 'lowres, bad anatomy, bad hands' or a specific inverse style.
- Settings: Keep CFG between 5.0 and 7.0; higher values can cause color burn.
- Sampler: Works well with DPM++ 2M SDE.
"""
"Resolution Set" = "anime_xl"
"CFG Scale" = 6.0
"Sampling Steps" = 28
"Sampling Method" = "DPM++ 2M SDE"
["flux1-dev"]
description = "State-of-the-art general purpose model, takes natural language."
guide = """
Detailed Prompting Guide:
- Style: Natural Language. Describe the scene as you would to a human.
- Negatives: Generally doesn't require a negative prompt.
- Settings: Works best with higher resolutions (1024+).
- Sampler: Flux Realistic or Euler.
"""
"Resolution Set" = "flux_standard"
"CFG Scale" = 3.5
"Sampling Steps" = 20
"Sampling Method" = "Flux Realistic"

15
resolution_presets.toml Normal file
View File

@ -0,0 +1,15 @@
# Resolution Presets
# Format: [preset_name]
# width and height are mandatory.
[anime_xl]
square = { width = 1024, height = 1024 }
portrait = { width = 832, height = 1216 }
landscape = { width = 1216, height = 832 }
widescreen = { width = 1536, height = 640 }
[flux_standard]
square = { width = 1024, height = 1024 }
portrait = { width = 896, height = 1152 }
landscape = { width = 1152, height = 896 }
widescreen = { width = 1344, height = 768 }

View File

@ -23,6 +23,8 @@ from .list_directory import handle as list_directory_details_handler
from .preview_image import handle as preview_image_handler from .preview_image import handle as preview_image_handler
from .get_text_context import handle as get_text_context_handler from .get_text_context import handle as get_text_context_handler
from .wikipedia import handle as wikipedia_handler from .wikipedia import handle as wikipedia_handler
from .get_model_info import handle as get_model_info_handler
from .generate_image import handle as generate_image_handler
# Central registry of all available tools # Central registry of all available tools
TOOL_REGISTRY = [ TOOL_REGISTRY = [
@ -124,4 +126,35 @@ TOOL_REGISTRY = [
}, },
handler=wikipedia_handler handler=wikipedia_handler
), ),
Tool(
name="get_model_info",
description="Returns a list of available models and their short descriptions. If a specific model name is provided, it returns the comprehensive prompting guide, available resolution presets, and active configuration tips for that model. CRITICAL: You must call this for any model you are unfamiliar with, as prompting styles vary wildly (e.g., tag-based vs. natural language) and using the wrong style will result in poor image quality.",
schema={
"type": "object",
"properties": {
"model_name": {"type": "string", "description": "The name of the model to get detailed info for. Leave empty to list all available models."}
},
"required": [],
},
handler=get_model_info_handler
),
Tool(
name="generate_image",
description="Generates an image using the specified model and parameters. Note: To ensure high quality, verify the model's prompting requirements via get_model_info before calling this tool.",
schema={
"type": "object",
"properties": {
"model_name": {"type": "string", "description": "The name of the model to use. Required."},
"prompt": {"type": "string", "description": "The prompt for the image. Required."},
"negative_prompt": {"type": "string", "description": "The negative prompt to exclude unwanted elements."},
"resolution_preset": {"type": "string", "description": "A named resolution preset (e.g., 'square', 'portrait'). Available options depend on the model."},
"steps": {"type": "integer", "description": "Number of sampling steps."},
"cfg_scale": {"type": "number", "description": "CFG scale for prompt adherence."},
"seed": {"type": "number", "description": "Random seed for reproducibility. Use -1 for random."},
"sampler": {"type": "string", "description": "The sampling method to use."},
},
"required": ["model_name", "prompt"],
},
handler=generate_image_handler
),
] ]

167
tools/generate_image.py Normal file
View File

@ -0,0 +1,167 @@
import requests
import tomllib
import os
import base64
import config
from typing import Any, Dict, List
from tools.utils import ToolError
# Paths to config files
MODEL_PRESETS_PATH = os.path.join(os.path.dirname(os.path.dirname(__file__)), "model_presets.toml")
RES_PRESETS_PATH = os.path.join(os.path.dirname(os.path.dirname(__file__)), "resolution_presets.toml")
SD_URL = "http://127.0.0.1:7860"
def load_toml(path: str) -> Dict[str, Any]:
try:
with open(path, "rb") as f:
return tomllib.load(f)
except Exception as e:
raise ToolError(f"Failed to load config file {path}: {str(e)}")
async def handle(args: Dict[str, Any]) -> List[Dict[str, Any]]:
"""
Generates an image using the specified model and parameters.
"""
# 1. REQUIRED ARGUMENTS
prompt = args.get("prompt")
if not prompt:
raise ToolError("The 'prompt' argument is required.")
model_name = args.get("model_name")
if not model_name:
raise ToolError("The 'model_name' argument is required. Use get_model_info to see available models.")
# 2. FETCH CURRENT SERVER DEFAULTS & MAP LABELS
try:
info_resp = requests.get(f"{SD_URL}/info", timeout=10)
info_resp.raise_for_status()
info_data = info_resp.json()
params_info = info_data["named_endpoints"]["/txt2img"]["parameters"]
except Exception as e:
raise ToolError(f"Failed to connect to Stable Diffusion server: {str(e)}")
# Build the label-to-index map
label_map = {}
label_counts = {}
payload = [p["parameter_default"] for p in params_info]
for idx, p in enumerate(params_info):
label = p["label"]
if not label or label.startswith("parameter_"):
label = p["parameter_name"]
if label in label_counts:
label_counts[label] += 1
mapped_label = f"{label} [{label_counts[label]}]"
else:
label_counts[label] = 0
mapped_label = label
label_map[mapped_label] = idx
# 3. LOAD CONFIGS
models_cfg = load_toml(MODEL_PRESETS_PATH)
res_cfg = load_toml(RES_PRESETS_PATH)
if model_name not in models_cfg:
raise ToolError(f"Model '{model_name}' not found in presets. Available: {', '.join(models_cfg.keys())}")
model_preset = models_cfg[model_name]
# 4. MERGE PIPELINE
for key, value in model_preset.items():
if key in ["description", "guide", "Resolution Set"]:
continue
if key in label_map:
payload[label_map[key]] = value
elif key.startswith("param_"):
try:
idx = int(key.replace("param_", ""))
if 0 <= idx < len(payload):
payload[idx] = value
except ValueError:
pass
res_preset_name = args.get("resolution_preset")
if res_preset_name:
res_set_name = model_preset.get("Resolution Set")
if res_set_name and res_set_name in res_cfg:
res_set = res_cfg[res_set_name]
if res_preset_name in res_set:
res_vals = res_set[res_preset_name]
if "Width" in label_map:
payload[label_map["Width"]] = res_vals["width"]
if "Height" in label_map:
payload[label_map["Height"]] = res_vals["height"]
else:
raise ToolError(f"Resolution preset '{res_preset_name}' not found for this model. Available: {', '.join(res_set.keys())}")
else:
raise ToolError(f"No resolution set configured for model '{model_name}'.")
overrides = {
"prompt": "Prompt",
"negative_prompt": "Negative Prompt",
"steps": "Sampling Steps",
"cfg_scale": "CFG Scale",
"seed": "Seed",
"sampler": "Sampling Method"
}
for arg_key, label in overrides.items():
if arg_key in args and label in label_map:
payload[label_map[label]] = args[arg_key]
if "Prompt" in label_map:
payload[label_map["Prompt"]] = prompt
for _ in range(7):
payload.insert(39, None)
# 6. EXECUTE GENERATION
try:
gen_payload = {"data": payload}
gen_resp = requests.post(f"{SD_URL}/api/txt2img", json=gen_payload, timeout=300)
gen_resp.raise_for_status()
res_data = gen_resp.json()
gallery_data = res_data["data"][0]["value"]
if not gallery_data:
raise ToolError("Server returned successfully but no image was generated.")
sd_img_url = gallery_data[0]["image"]["url"]
# Extract raw path from SD URL (e.g. http://.../file=/tmp/gradio/abc.png -> /tmp/gradio/abc.png)
if "/file=" in sd_img_url:
raw_path = sd_img_url.split("/file=")[1]
else:
# Fallback if the URL format changes
raise ToolError(f"Could not extract file path from SD URL: {sd_img_url}")
# Construct proxy URL through our MCP server
# We strip leading slash from raw_path to avoid double slashes in the proxy URL if we want,
# but the current endpoint handles it.
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.raise_for_status()
b64_data = base64.b64encode(img_resp.content).decode('utf-8')
return [
{
"type": "text",
"text": f"Image successfully generated using model '{model_name}'.\n\nTo show the image to the user, you can include this markdown link in your response: ![Generated Image]({proxy_url})"
},
{
"type": "image",
"data": b64_data,
"mimeType": "image/png"
}
]
except requests.exceptions.HTTPError as e:
raise ToolError(f"Server error during generation: {str(e)}")
except Exception as e:
raise ToolError(f"Unexpected error during generation: {str(e)}")

70
tools/get_model_info.py Normal file
View File

@ -0,0 +1,70 @@
import tomllib
import os
from typing import Any, Dict, Union, List
from tools.utils import ToolError
# Paths to config files
MODEL_PRESETS_PATH = os.path.join(os.path.dirname(os.path.dirname(__file__)), "model_presets.toml")
RES_PRESETS_PATH = os.path.join(os.path.dirname(os.path.dirname(__file__)), "resolution_presets.toml")
def load_toml(path: str) -> Dict[str, Any]:
try:
with open(path, "rb") as f:
return tomllib.load(f)
except Exception as e:
raise ToolError(f"Failed to load config file {path}: {str(e)}")
async def handle(args: Dict[str, Any]) -> Dict[str, Any]:
"""
Returns information about available models or detailed guides for a specific model.
"""
model_name = args.get("model_name")
models = load_toml(MODEL_PRESETS_PATH)
res_presets = load_toml(RES_PRESETS_PATH)
if not model_name:
# Return a catalog of all models
catalog = []
for name, data in models.items():
catalog.append({
"name": name,
"description": data.get("description", "No description provided.")
})
return {
"text": "Available 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."
}
if model_name not in models:
raise ToolError(f"Model '{model_name}' not found in presets. Available models: {', '.join(models.keys())}")
model_data = models[model_name]
res_set_name = model_data.get("Resolution Set")
# Get available resolution presets for this model
res_options = []
if res_set_name and res_set_name in res_presets:
res_options = list(res_presets[res_set_name].keys())
else:
res_options = ["Default (1024x1024)"]
guide = model_data.get("guide", "No detailed guide available.")
description = model_data.get("description", "")
# Filter out internal config keys to show only the human-friendly presets
presets_to_show = {k: v for k, v in model_data.items() if k not in ["description", "guide", "Resolution Set"]}
preset_text = "\n".join([f"- {k}: {v}" for k, v in presets_to_show.items()])
res_text = ", ".join(res_options)
return {
"text": (
f"Model: {model_name}\n"
f"Description: {description}\n\n"
f"--- Prompting Guide ---\n{guide}\n\n"
f"--- Active Presets ---\n{preset_text}\n\n"
f"--- Available Resolution Presets ---\n{res_text}\n\n"
f"Use 'resolution_preset' in generate_image to choose one of these."
)
}