980 lines
34 KiB
Python
980 lines
34 KiB
Python
#!/usr/bin/env python3
|
||
"""
|
||
MCP server for generating pixel-art sprites using SDXL + LCM + pixel-art-xl LoRA.
|
||
|
||
Tools:
|
||
- generate_sprite: Generate a single pixel-art sprite
|
||
- batch_generate: Generate multiple sprites in one call
|
||
- rate_sprite: Rate a generated sprite (1-5 stars) with optional feedback
|
||
- get_reference_sprites: Get highly-rated reference sprites for a prompt
|
||
- list_sprites: List sprites in the feedback DB (all, unrated, or top-rated)
|
||
- db_stats: Get feedback database statistics
|
||
|
||
Model is loaded lazily on first call (~12s), then stays in VRAM for speed.
|
||
Background is removed post-generation to produce transparent PNG.
|
||
Every generated sprite is automatically saved to the feedback DB (unrated).
|
||
"""
|
||
|
||
import os
|
||
import sys
|
||
import time
|
||
import base64
|
||
import json
|
||
from io import BytesIO
|
||
from typing import Optional
|
||
from urllib.request import urlopen
|
||
|
||
import numpy as np
|
||
from PIL import Image, PngImagePlugin
|
||
from mcp.server.fastmcp import FastMCP
|
||
|
||
from feedback import FeedbackDB
|
||
from polza import PolzaClient, PolzaError
|
||
from pixellab import PixelLabClient
|
||
|
||
|
||
def _load_dotenv(path: str) -> None:
|
||
"""Load a minimal KEY=value .env file without overriding real environment."""
|
||
if not os.path.isfile(path):
|
||
return
|
||
with open(path, encoding="utf-8") as env_file:
|
||
for raw_line in env_file:
|
||
line = raw_line.strip()
|
||
if not line or line.startswith("#") or "=" not in line:
|
||
continue
|
||
key, value = line.split("=", 1)
|
||
key, value = key.strip(), value.strip()
|
||
if value[:1] == value[-1:] and value[:1] in {"'", '"'}:
|
||
value = value[1:-1]
|
||
if key:
|
||
os.environ.setdefault(key, value)
|
||
|
||
|
||
# Paths — models live in a shared location
|
||
BASE_DIR = os.path.dirname(os.path.abspath(__file__))
|
||
_load_dotenv(os.path.join(BASE_DIR, ".env"))
|
||
MODEL_DIR = os.environ.get(
|
||
"IMAGEGEN_MODEL_DIR",
|
||
os.path.join(os.path.expanduser("~"), "models", "sdxl-base"),
|
||
)
|
||
LORA_DIR = os.environ.get(
|
||
"IMAGEGEN_LORA_DIR",
|
||
os.path.join(os.path.expanduser("~"), "models", "pixel-art-xl"),
|
||
)
|
||
LCM_LORA_DIR = os.environ.get(
|
||
"IMAGEGEN_LCM_LORA_DIR",
|
||
os.path.join(os.path.expanduser("~"), "models", "lcm-lora-sdxl"),
|
||
)
|
||
OUTPUT_DIR = os.environ.get("IMAGEGEN_OUTPUT_DIR", os.path.join(BASE_DIR, "output"))
|
||
DB_PATH = os.environ.get("IMAGEGEN_DB_PATH", os.path.join(BASE_DIR, "feedback.db"))
|
||
POLZA_MODEL = os.environ.get("POLZA_IMAGE_MODEL", "openai/gpt-image-1.5")
|
||
|
||
# LoRA scales — pixel-art-xl needs 1.2, LCM needs 1.0
|
||
PIXEL_LORA_SCALE = 1.2
|
||
LCM_LORA_SCALE = 1.0
|
||
|
||
# Negative prompt to improve quality
|
||
NEGATIVE_PROMPT = "3d render, realistic, detailed, noise, artifacts, blurry, text"
|
||
|
||
# Global state — model loaded lazily
|
||
_pipe = None
|
||
_device = None
|
||
_db: Optional[FeedbackDB] = None
|
||
|
||
|
||
def _get_db() -> FeedbackDB:
|
||
global _db
|
||
if _db is None:
|
||
_db = FeedbackDB.open(DB_PATH)
|
||
sys.stderr.write(f"[pixel-art] Feedback DB: {DB_PATH}\n")
|
||
return _db
|
||
|
||
|
||
def _get_device():
|
||
global _device
|
||
if _device is None:
|
||
import torch
|
||
|
||
if torch.cuda.is_available():
|
||
_device = "cuda"
|
||
else:
|
||
_device = "cpu"
|
||
sys.stderr.write(
|
||
"[pixel-art] WARNING: CUDA not available, using CPU (very slow)\n"
|
||
)
|
||
return _device
|
||
|
||
|
||
def _load_model():
|
||
global _pipe
|
||
if _pipe is not None:
|
||
return _pipe
|
||
|
||
sys.stderr.write(
|
||
"[pixel-art] Loading SDXL + LCM + pixel-art-xl LoRA (first call)...\n"
|
||
)
|
||
t0 = time.time()
|
||
|
||
import torch
|
||
from diffusers import StableDiffusionXLPipeline, LCMScheduler
|
||
|
||
_pipe = StableDiffusionXLPipeline.from_pretrained(
|
||
MODEL_DIR,
|
||
torch_dtype=torch.float16,
|
||
use_safetensors=True,
|
||
variant="fp16",
|
||
)
|
||
_pipe.scheduler = LCMScheduler.from_config(_pipe.scheduler.config)
|
||
|
||
_pipe.load_lora_weights(LCM_LORA_DIR, adapter_name="lcm")
|
||
_pipe.load_lora_weights(LORA_DIR, adapter_name="pixel")
|
||
_pipe.set_adapters(
|
||
["lcm", "pixel"], adapter_weights=[LCM_LORA_SCALE, PIXEL_LORA_SCALE]
|
||
)
|
||
|
||
if _get_device() == "cuda":
|
||
_pipe.to("cuda")
|
||
else:
|
||
_pipe.to(_get_device())
|
||
|
||
elapsed = time.time() - t0
|
||
sys.stderr.write(f"[pixel-art] Model loaded in {elapsed:.1f}s\n")
|
||
return _pipe
|
||
|
||
|
||
def _build_prompt(user_prompt: str) -> str:
|
||
return f"pixel art, {user_prompt}, simple, flat colors, game asset"
|
||
|
||
|
||
def _generate(
|
||
pipe, prompt: str, seed: Optional[int], width: int, height: int, steps: int
|
||
):
|
||
import torch
|
||
|
||
generator = None
|
||
if seed is not None:
|
||
generator = torch.Generator(device="cuda").manual_seed(seed)
|
||
|
||
image = pipe(
|
||
prompt=prompt,
|
||
negative_prompt=NEGATIVE_PROMPT,
|
||
num_inference_steps=steps,
|
||
guidance_scale=1.5,
|
||
height=height,
|
||
width=width,
|
||
generator=generator,
|
||
).images[0]
|
||
|
||
return image
|
||
|
||
|
||
def _remove_background(image: Image.Image, threshold: int = 30) -> Image.Image:
|
||
"""Remove background using flood-fill from edges.
|
||
|
||
Detects the border color, then flood-fills from all edge pixels,
|
||
removing any pixel that is within threshold of the border color AND
|
||
connected to the border. Interior pixels with similar colors (e.g.
|
||
highlights on armor) are preserved because they are not connected
|
||
to the border through similar-colored pixels.
|
||
"""
|
||
from collections import deque
|
||
|
||
rgb = image.convert("RGB")
|
||
arr = np.array(rgb).astype(int)
|
||
h, w = arr.shape[:2]
|
||
|
||
# Sample border colors from all 4 edges
|
||
border_colors = []
|
||
for x in range(w):
|
||
border_colors.append(arr[0, x])
|
||
border_colors.append(arr[h - 1, x])
|
||
for y in range(h):
|
||
border_colors.append(arr[y, 0])
|
||
border_colors.append(arr[y, w - 1])
|
||
|
||
border_colors = np.array(border_colors)
|
||
bg_color = np.median(border_colors, axis=0).astype(int)
|
||
|
||
alpha = np.full((h, w), 255, dtype=np.uint8)
|
||
visited = np.zeros((h, w), dtype=bool)
|
||
queue = deque()
|
||
|
||
# Seed from all border pixels
|
||
for x in range(w):
|
||
for y in [0, h - 1]:
|
||
if not visited[y, x]:
|
||
queue.append((y, x))
|
||
visited[y, x] = True
|
||
for y in range(h):
|
||
for x in [0, w - 1]:
|
||
if not visited[y, x]:
|
||
queue.append((y, x))
|
||
visited[y, x] = True
|
||
|
||
# BFS flood-fill: remove pixels close to bg_color that are connected to border
|
||
while queue:
|
||
y, x = queue.popleft()
|
||
dist = np.abs(arr[y, x] - bg_color).sum()
|
||
if dist > threshold * 3:
|
||
continue
|
||
alpha[y, x] = 0
|
||
|
||
for dy, dx in [(-1, 0), (1, 0), (0, -1), (0, 1)]:
|
||
ny, nx = y + dy, x + dx
|
||
if 0 <= ny < h and 0 <= nx < w and not visited[ny, nx]:
|
||
visited[ny, nx] = True
|
||
queue.append((ny, nx))
|
||
|
||
rgba = np.dstack([arr.astype(np.uint8), alpha])
|
||
return Image.fromarray(rgba, mode="RGBA")
|
||
|
||
|
||
def _pixelate(image: Image.Image, pixel_size: int = 8) -> Image.Image:
|
||
"""Downscale then upscale with NEAREST to create chunky pixel-art effect.
|
||
|
||
pixel_size=8 means each "pixel" in the result is an 8x8 block.
|
||
pixel_size=0 returns the original image unchanged.
|
||
"""
|
||
if pixel_size <= 0:
|
||
return image
|
||
w, h = image.size
|
||
small = image.resize((w // pixel_size, h // pixel_size), Image.LANCZOS)
|
||
return small.resize((w, h), Image.NEAREST)
|
||
|
||
|
||
def _ensure_dir(path: str):
|
||
dir_path = os.path.dirname(path)
|
||
if dir_path:
|
||
os.makedirs(dir_path, exist_ok=True)
|
||
|
||
|
||
def _save_sprite(image: Image.Image, path: str, prompt: str, metadata: dict) -> None:
|
||
"""Save a PNG with portable, machine-readable generation provenance."""
|
||
_ensure_dir(path)
|
||
png_info = PngImagePlugin.PngInfo()
|
||
png_info.add_text("prompt", prompt)
|
||
png_info.add_text("imagen", json.dumps(metadata, ensure_ascii=False))
|
||
image.save(path, format="PNG", pnginfo=png_info)
|
||
|
||
|
||
def _output_paths(output_path: str, count: int) -> list[str]:
|
||
"""Keep a requested filename for one result; suffix variants for many."""
|
||
if not os.path.isabs(output_path):
|
||
output_path = os.path.join(OUTPUT_DIR, output_path)
|
||
if count == 1:
|
||
return [output_path]
|
||
stem, extension = os.path.splitext(output_path)
|
||
extension = extension or ".png"
|
||
return [f"{stem}_{index:02d}{extension}" for index in range(1, count + 1)]
|
||
|
||
|
||
def _image_from_source(source: str) -> Image.Image:
|
||
if source.startswith("data:"):
|
||
source = source.split(",", 1)[-1]
|
||
if source.startswith(("http://", "https://")):
|
||
with urlopen(source, timeout=60) as response:
|
||
data = response.read()
|
||
else:
|
||
data = base64.b64decode(source)
|
||
return Image.open(BytesIO(data)).convert("RGBA")
|
||
|
||
|
||
def _polza_generate(
|
||
prompt: str,
|
||
output_path: str,
|
||
*,
|
||
count: int = 1,
|
||
model: Optional[str] = None,
|
||
reference_images: Optional[list[str]] = None,
|
||
aspect_ratio: Optional[str] = None,
|
||
seed: Optional[int] = None,
|
||
quality: Optional[str] = None,
|
||
remove_bg: bool = True,
|
||
pixel_size: int = 4,
|
||
wait: bool = True,
|
||
) -> list[dict]:
|
||
"""Generate and persist image variants through Polza Media API."""
|
||
client = PolzaClient()
|
||
full_prompt = _build_prompt(prompt)
|
||
started = time.time()
|
||
selected_model = model or POLZA_MODEL
|
||
response = client.create_image(
|
||
model=selected_model,
|
||
prompt=full_prompt,
|
||
reference_images=reference_images,
|
||
count=count,
|
||
aspect_ratio=aspect_ratio,
|
||
seed=seed,
|
||
quality=quality,
|
||
# The transparent-background option is specific to GPT Image. Other
|
||
# models still get transparent PNGs through local post-processing.
|
||
background="transparent" if remove_bg and "gpt-image" in selected_model else None,
|
||
wait=wait,
|
||
)
|
||
task_id = response.get("id")
|
||
state = response.get("status", "").lower()
|
||
if wait and task_id and state in {"pending", "queued", "processing", "running"}:
|
||
response = client.wait_for_completion(task_id)
|
||
|
||
sources = client.image_sources(response)
|
||
if not sources:
|
||
return [{
|
||
"generation_id": task_id,
|
||
"status": response.get("status", "pending"),
|
||
"provider": "polza",
|
||
"model": response.get("model", model or POLZA_MODEL),
|
||
"message": "Generation is still running; call get_generation_status later.",
|
||
}]
|
||
|
||
paths = _output_paths(output_path, len(sources))
|
||
results = []
|
||
for index, (source, path) in enumerate(zip(sources, paths), start=1):
|
||
_ensure_dir(path)
|
||
image = _image_from_source(source)
|
||
if pixel_size > 0:
|
||
image = _pixelate(image, pixel_size)
|
||
if remove_bg:
|
||
image = _remove_background(image)
|
||
_save_sprite(image, path, prompt, {"provider": "polza", "full_prompt": full_prompt,
|
||
"generation_id": task_id, "model": response.get("model", model or POLZA_MODEL)})
|
||
params = {
|
||
"provider": "polza",
|
||
"generation_id": task_id,
|
||
"model": response.get("model", model or POLZA_MODEL),
|
||
"seed": seed,
|
||
"reference_images": reference_images or [],
|
||
"aspect_ratio": aspect_ratio,
|
||
"quality": quality,
|
||
"remove_bg": remove_bg,
|
||
"pixel_size": pixel_size,
|
||
"full_prompt": full_prompt,
|
||
"usage": response.get("usage"),
|
||
"warnings": response.get("warnings", []),
|
||
}
|
||
entry_id = _get_db().add(prompt=prompt, params=params, image_path=path)
|
||
results.append({
|
||
"output_path": path,
|
||
"source_url": source if source.startswith(("http://", "https://")) else None,
|
||
"variant": index,
|
||
"db_id": entry_id,
|
||
"generation_id": task_id,
|
||
"generation_time": f"{time.time() - started:.1f}s",
|
||
"prompt": full_prompt,
|
||
"provider": "polza",
|
||
"model": response.get("model", model or POLZA_MODEL),
|
||
"usage": response.get("usage"),
|
||
"warnings": response.get("warnings", []),
|
||
"rated": False,
|
||
})
|
||
return results
|
||
|
||
|
||
def _pixellab_generate(
|
||
prompt: str, output_path: str, *, count: int = 1, reference_images: Optional[list[str]] = None,
|
||
seed: Optional[int] = None, width: int = 128, height: int = 128, remove_bg: bool = True,
|
||
pixel_size: int = 0, wait: bool = True,
|
||
) -> list[dict]:
|
||
"""Generate native pixel-art sprites through PixelLab v2."""
|
||
client = PixelLabClient()
|
||
full_prompt = _build_prompt(prompt)
|
||
responses = [client.create_image(full_prompt, width, height, seed=seed, reference_images=reference_images,
|
||
no_background=remove_bg) for _ in range(count)]
|
||
completed = []
|
||
for response in responses:
|
||
job_id = response.get("background_job_id")
|
||
if wait and job_id:
|
||
response = client.wait_for_completion(job_id)
|
||
completed.append((job_id, response))
|
||
source_records = [(job_id, response, source) for job_id, response in completed
|
||
for source in client.image_sources(response)]
|
||
if not source_records:
|
||
return [{"generation_id": job_id, "status": response.get("status", "processing"),
|
||
"provider": "pixellab", "usage": response.get("usage")}
|
||
for job_id, response in completed]
|
||
results = []
|
||
for index, ((job_id, response, source), path) in enumerate(
|
||
zip(source_records, _output_paths(output_path, len(source_records))), start=1
|
||
):
|
||
_ensure_dir(path)
|
||
image = _image_from_source(source)
|
||
if pixel_size > 0:
|
||
image = _pixelate(image, pixel_size)
|
||
if remove_bg:
|
||
image = _remove_background(image)
|
||
_save_sprite(image, path, prompt, {"provider": "pixellab", "full_prompt": full_prompt,
|
||
"generation_id": job_id, "width": width, "height": height})
|
||
db_id = _get_db().add(prompt=prompt, image_path=path, params={
|
||
"provider": "pixellab", "generation_id": job_id, "reference_images": reference_images or [],
|
||
"seed": seed, "width": width, "height": height, "remove_bg": remove_bg,
|
||
"pixel_size": pixel_size, "full_prompt": full_prompt, "usage": response.get("usage"),
|
||
})
|
||
results.append({"output_path": path, "db_id": db_id, "generation_id": job_id, "variant": index,
|
||
"provider": "pixellab", "usage": response.get("usage"), "rated": False})
|
||
return results
|
||
|
||
|
||
# Instructions are part of the MCP initialization response, so every connected
|
||
# agent receives the asset-generation rules before it chooses a tool.
|
||
MCP_INSTRUCTIONS = """
|
||
Generate game assets deliberately and keep the dataset reusable.
|
||
|
||
Style workflow:
|
||
1. Call get_project_style_guide first. Its registered references live in the
|
||
server database; pass 1–4 selected paths as reference_images to
|
||
provider='pixellab' for a coherent style family. get_reference_sprites is
|
||
useful for additional rated examples.
|
||
2. Generate one semantic asset role per prompt (for example only a missile
|
||
corvette). Never put a list of roles in one prompt when individual metadata
|
||
must identify each image: all outputs of a multi-image batch inherit the
|
||
same prompt and cannot reliably be labelled afterwards.
|
||
3. State camera, orientation, silhouette, palette, background, and exclusions
|
||
explicitly. For top-down fleets, say: 'orthographic top-down; nose at 12
|
||
o'clock; engines at 6 o'clock'. Add exact vertical mirror symmetry only
|
||
when it is artistically required.
|
||
|
||
Pixel-art rules:
|
||
- For pixel-perfect assets use PixelLab at a native square size (normally
|
||
64x64 or 128x128), transparent background, pixel_size=0, and request no
|
||
anti-aliasing, blur, sub-pixel shading, or semi-transparent edges.
|
||
- In HTML/game clients display these files at an integer multiple with
|
||
image-rendering: pixelated; never use arbitrary CSS scaling.
|
||
- Use provider='pixellab' for native game sprites and style references.
|
||
Use provider='polza' for broader concept art or backgrounds. Local is an
|
||
optional Diffusers/GPU provider.
|
||
|
||
Operations and provenance:
|
||
- Use a semantic output_path and preserve every returned db_id. Finished PNGs
|
||
embed 'prompt' and 'imagen' JSON metadata (provider, full prompt and
|
||
generation settings); the feedback database stores the same provenance.
|
||
- generate_images is for variants of one precise asset. batch_generate is for
|
||
several independent, precisely described assets. For async work use wait=false
|
||
and get_generation_status with the matching provider.
|
||
- Never expose, request, or store API keys in prompts, output paths, feedback,
|
||
or image metadata.
|
||
""".strip()
|
||
|
||
|
||
# Create MCP server
|
||
mcp = FastMCP("pixel-art", instructions=MCP_INSTRUCTIONS)
|
||
|
||
|
||
@mcp.tool()
|
||
def generate_sprite(
|
||
prompt: str,
|
||
output_path: str,
|
||
seed: Optional[int] = None,
|
||
width: int = 512,
|
||
height: int = 512,
|
||
steps: int = 8,
|
||
remove_bg: bool = True,
|
||
pixel_size: int = 4,
|
||
provider: str = "local",
|
||
model: Optional[str] = None,
|
||
reference_images: Optional[list[str]] = None,
|
||
aspect_ratio: Optional[str] = None,
|
||
quality: Optional[str] = None,
|
||
) -> dict:
|
||
"""Generate one precisely described pixel-art sprite and save it as PNG.
|
||
|
||
For a reusable game dataset, describe only one semantic role per call and
|
||
state view/orientation/palette explicitly. Use provider='pixellab' plus
|
||
1–4 reference_images for style-consistent native pixel art. Each PNG
|
||
embeds the prompt and generation provenance in its metadata.
|
||
|
||
Args:
|
||
prompt: Description of the sprite (e.g. "a crystal warrior with geometric armor")
|
||
output_path: Where to save the PNG file (relative to output dir or absolute)
|
||
seed: Optional seed for reproducibility
|
||
width: Image width in pixels (default 512)
|
||
height: Image height in pixels (default 512)
|
||
steps: Inference steps (default 4, FLUX.2-klein is distilled)
|
||
remove_bg: Remove background and make transparent (default True)
|
||
pixel_size: Size of each pixel block for pixel-art effect (default 4, 0=off)
|
||
provider: "local" (SDXL) or "polza" (cloud image models)
|
||
model: Polza model ID; defaults to POLZA_IMAGE_MODEL
|
||
reference_images: Style references: HTTPS URLs, data URIs, or local image paths
|
||
aspect_ratio: Cloud aspect ratio such as "1:1" or "16:9"
|
||
quality: Cloud model quality setting
|
||
|
||
Returns:
|
||
Dict with output_path, seed_used, generation_time, prompt, size.
|
||
"""
|
||
if provider == "polza":
|
||
return _polza_generate(
|
||
prompt,
|
||
output_path,
|
||
model=model,
|
||
reference_images=reference_images,
|
||
aspect_ratio=aspect_ratio,
|
||
seed=seed,
|
||
quality=quality,
|
||
remove_bg=remove_bg,
|
||
pixel_size=pixel_size,
|
||
)[0]
|
||
if provider == "pixellab":
|
||
return _pixellab_generate(prompt, output_path, reference_images=reference_images, seed=seed,
|
||
width=width, height=height, remove_bg=remove_bg, pixel_size=pixel_size)[0]
|
||
if provider != "local":
|
||
raise ValueError("provider must be 'local', 'polza', or 'pixellab'")
|
||
if reference_images:
|
||
raise ValueError("reference_images require provider='polza'")
|
||
|
||
pipe = _load_model()
|
||
full_prompt = _build_prompt(prompt)
|
||
|
||
if not os.path.isabs(output_path):
|
||
output_path = os.path.join(OUTPUT_DIR, output_path)
|
||
|
||
_ensure_dir(output_path)
|
||
|
||
t0 = time.time()
|
||
image = _generate(pipe, full_prompt, seed, width, height, steps)
|
||
|
||
if pixel_size > 0:
|
||
image = _pixelate(image, pixel_size)
|
||
|
||
if remove_bg:
|
||
image = _remove_background(image)
|
||
|
||
_save_sprite(image, output_path, prompt, {"provider": "local", "full_prompt": full_prompt,
|
||
"seed": seed, "width": width, "height": height, "steps": steps})
|
||
elapsed = time.time() - t0
|
||
|
||
db = _get_db()
|
||
entry_id = db.add(
|
||
prompt=prompt,
|
||
params={
|
||
"seed": seed,
|
||
"width": width,
|
||
"height": height,
|
||
"steps": steps,
|
||
"remove_bg": remove_bg,
|
||
"pixel_size": pixel_size,
|
||
"model": MODEL_DIR,
|
||
"lora": LORA_DIR,
|
||
"lcm": LCM_LORA_DIR,
|
||
"lora_scale": PIXEL_LORA_SCALE,
|
||
"lcm_scale": LCM_LORA_SCALE,
|
||
"negative_prompt": NEGATIVE_PROMPT,
|
||
"guidance_scale": 1.5,
|
||
"scheduler": "LCMScheduler",
|
||
"full_prompt": full_prompt,
|
||
},
|
||
image_path=output_path,
|
||
)
|
||
|
||
return {
|
||
"output_path": output_path,
|
||
"seed_used": seed,
|
||
"generation_time": f"{elapsed:.1f}s",
|
||
"prompt": full_prompt,
|
||
"size": f"{width}x{height}",
|
||
"transparent": remove_bg,
|
||
"pixel_size": pixel_size,
|
||
"db_id": entry_id,
|
||
"rated": False,
|
||
}
|
||
|
||
|
||
@mcp.tool()
|
||
def batch_generate(
|
||
specs: list[dict],
|
||
) -> list[dict]:
|
||
"""Generate multiple independent, precisely described sprites in one call.
|
||
|
||
Each spec must represent one asset role. Do not ask one spec for mixed
|
||
categories if you need per-image semantic metadata: every result inherits
|
||
that spec's single prompt. For a game style family, give every PixelLab spec
|
||
the same selected reference_images and native 64x64 or 128x128 dimensions.
|
||
|
||
Args:
|
||
specs: List of dicts, each with:
|
||
- prompt: str (required) — sprite description
|
||
- output_path: str (required) — PNG save path
|
||
- seed: int (optional)
|
||
- width: int (optional, default 512)
|
||
- height: int (optional, default 512)
|
||
- steps: int (optional, default 8)
|
||
- remove_bg: bool (optional, default True)
|
||
- pixel_size: int (optional, default 4, 0=off)
|
||
|
||
Returns:
|
||
List of dicts with output_path, seed_used, generation_time, prompt, size, transparent.
|
||
"""
|
||
pipe = None
|
||
results = []
|
||
|
||
for spec in specs:
|
||
prompt = spec["prompt"]
|
||
output_path = spec["output_path"]
|
||
seed = spec.get("seed")
|
||
width = spec.get("width", 512)
|
||
height = spec.get("height", 512)
|
||
steps = spec.get("steps", 8)
|
||
remove_bg = spec.get("remove_bg", True)
|
||
pixel_size = spec.get("pixel_size", 4)
|
||
provider = spec.get("provider", "local")
|
||
model = spec.get("model")
|
||
reference_images = spec.get("reference_images")
|
||
aspect_ratio = spec.get("aspect_ratio")
|
||
quality = spec.get("quality")
|
||
|
||
if provider == "polza":
|
||
results.extend(
|
||
_polza_generate(
|
||
prompt,
|
||
output_path,
|
||
count=spec.get("count", 1),
|
||
model=model,
|
||
reference_images=reference_images,
|
||
aspect_ratio=aspect_ratio,
|
||
seed=seed,
|
||
quality=quality,
|
||
remove_bg=remove_bg,
|
||
pixel_size=pixel_size,
|
||
wait=spec.get("wait", True),
|
||
)
|
||
)
|
||
continue
|
||
if provider == "pixellab":
|
||
results.extend(
|
||
_pixellab_generate(
|
||
prompt,
|
||
output_path,
|
||
count=spec.get("count", 1),
|
||
reference_images=reference_images,
|
||
seed=seed,
|
||
width=width,
|
||
height=height,
|
||
remove_bg=remove_bg,
|
||
pixel_size=pixel_size,
|
||
wait=spec.get("wait", True),
|
||
)
|
||
)
|
||
continue
|
||
if provider != "local":
|
||
raise ValueError("provider must be 'local', 'polza', or 'pixellab'")
|
||
if reference_images:
|
||
raise ValueError("reference_images require provider='polza'")
|
||
if pipe is None:
|
||
pipe = _load_model()
|
||
|
||
full_prompt = _build_prompt(prompt)
|
||
|
||
if not os.path.isabs(output_path):
|
||
output_path = os.path.join(OUTPUT_DIR, output_path)
|
||
|
||
_ensure_dir(output_path)
|
||
|
||
t0 = time.time()
|
||
image = _generate(pipe, full_prompt, seed, width, height, steps)
|
||
|
||
if pixel_size > 0:
|
||
image = _pixelate(image, pixel_size)
|
||
|
||
if remove_bg:
|
||
image = _remove_background(image)
|
||
|
||
_save_sprite(image, output_path, prompt, {"provider": "local", "full_prompt": full_prompt,
|
||
"seed": seed, "width": width, "height": height, "steps": steps})
|
||
elapsed = time.time() - t0
|
||
|
||
db = _get_db()
|
||
entry_id = db.add(
|
||
prompt=prompt,
|
||
params={
|
||
"seed": seed,
|
||
"width": width,
|
||
"height": height,
|
||
"steps": steps,
|
||
"remove_bg": remove_bg,
|
||
"pixel_size": pixel_size,
|
||
"model": MODEL_DIR,
|
||
"lora": LORA_DIR,
|
||
"lcm": LCM_LORA_DIR,
|
||
"lora_scale": PIXEL_LORA_SCALE,
|
||
"lcm_scale": LCM_LORA_SCALE,
|
||
"negative_prompt": NEGATIVE_PROMPT,
|
||
"guidance_scale": 1.5,
|
||
"scheduler": "LCMScheduler",
|
||
"full_prompt": full_prompt,
|
||
},
|
||
image_path=output_path,
|
||
)
|
||
|
||
results.append(
|
||
{
|
||
"output_path": output_path,
|
||
"seed_used": seed,
|
||
"generation_time": f"{elapsed:.1f}s",
|
||
"prompt": full_prompt,
|
||
"size": f"{width}x{height}",
|
||
"transparent": remove_bg,
|
||
"db_id": entry_id,
|
||
"rated": False,
|
||
}
|
||
)
|
||
|
||
return results
|
||
|
||
|
||
@mcp.tool()
|
||
def generate_images(
|
||
prompt: str,
|
||
output_path: str,
|
||
count: int = 1,
|
||
model: Optional[str] = None,
|
||
reference_images: Optional[list[str]] = None,
|
||
aspect_ratio: Optional[str] = None,
|
||
seed: Optional[int] = None,
|
||
quality: Optional[str] = None,
|
||
remove_bg: bool = True,
|
||
pixel_size: int = 4,
|
||
wait: bool = True,
|
||
provider: str = "polza",
|
||
) -> list[dict]:
|
||
"""Generate cloud-image variations of one precise asset request.
|
||
|
||
Use ``reference_images`` for a game's style guide, existing characters, or
|
||
tiles. Each value may be an HTTPS URL, data URI, or a path readable by this
|
||
MCP server. Results use output_path_01.png, output_path_02.png, etc.
|
||
When ``wait`` is false, the returned generation_id can be passed to
|
||
get_generation_status later. Use separate calls (or batch_generate specs)
|
||
for different roles such as scout, corvette, and freighter, so embedded
|
||
prompt metadata remains meaningful for every image. For pixel-perfect
|
||
assets choose provider='pixellab'; it uses a native 128x128 output here.
|
||
"""
|
||
if provider == "pixellab":
|
||
# PixelLab produces native pixel art; its practical size range is 16–512.
|
||
# aspect_ratio is not applicable because its API accepts explicit dimensions.
|
||
return _pixellab_generate(
|
||
prompt, output_path, count=count, reference_images=reference_images, seed=seed,
|
||
remove_bg=remove_bg, pixel_size=pixel_size, wait=wait,
|
||
)
|
||
if provider != "polza":
|
||
raise ValueError("provider must be 'polza' or 'pixellab'")
|
||
return _polza_generate(
|
||
prompt,
|
||
output_path,
|
||
count=count,
|
||
model=model,
|
||
reference_images=reference_images,
|
||
aspect_ratio=aspect_ratio,
|
||
seed=seed,
|
||
quality=quality,
|
||
remove_bg=remove_bg,
|
||
pixel_size=pixel_size,
|
||
wait=wait,
|
||
)
|
||
|
||
|
||
@mcp.tool()
|
||
def get_generation_status(generation_id: str, provider: str = "polza") -> dict:
|
||
"""Get a Polza or PixelLab cloud-generation status and finished sources."""
|
||
if provider == "pixellab":
|
||
status = PixelLabClient().get_status(generation_id)
|
||
return {
|
||
"generation_id": status.get("id", generation_id), "status": status.get("status"),
|
||
"provider": "pixellab", "image_sources": PixelLabClient.image_sources(status),
|
||
"usage": status.get("usage"), "error": status.get("error"),
|
||
}
|
||
if provider != "polza":
|
||
raise ValueError("provider must be 'polza' or 'pixellab'")
|
||
status = PolzaClient().get_status(generation_id)
|
||
return {
|
||
"generation_id": status.get("id", generation_id),
|
||
"status": status.get("status"),
|
||
"model": status.get("model"),
|
||
"image_sources": PolzaClient.image_sources(status),
|
||
"usage": status.get("usage"),
|
||
"warnings": status.get("warnings", []),
|
||
"error": status.get("error"),
|
||
}
|
||
|
||
|
||
@mcp.tool()
|
||
def rate_sprite(
|
||
db_id: str,
|
||
rating: int,
|
||
feedback: Optional[str] = None,
|
||
) -> dict:
|
||
"""Rate a generated sprite (1-5 stars) with optional feedback text.
|
||
|
||
Use this after reviewing a sprite to teach the system what looks good.
|
||
The AI uses high-rated sprites as reference when generating similar ones.
|
||
|
||
Args:
|
||
db_id: The ID returned by generate_sprite or batch_generate
|
||
rating: 1-5 stars (5 = excellent, 1 = terrible)
|
||
feedback: Optional text feedback (e.g. "great colors, bad proportions")
|
||
|
||
Returns:
|
||
Dict with db_id, rating, feedback, and status.
|
||
"""
|
||
db = _get_db()
|
||
db.update_rating(db_id, rating, feedback)
|
||
|
||
return {
|
||
"db_id": db_id,
|
||
"rating": rating,
|
||
"feedback": feedback,
|
||
"status": "saved",
|
||
}
|
||
|
||
|
||
@mcp.tool()
|
||
def get_reference_sprites(
|
||
prompt: str,
|
||
limit: int = 5,
|
||
min_rating: int = 4,
|
||
) -> list[dict]:
|
||
"""Get highly-rated reference sprites from the feedback DB for a given prompt.
|
||
|
||
Use these as examples when generating similar sprites to improve quality.
|
||
Returns sprites with similar prompt keywords that have been rated >= min_rating.
|
||
|
||
Args:
|
||
prompt: The prompt to search for (e.g. "knight", "crystal warrior")
|
||
limit: Max number of results (default 5)
|
||
min_rating: Minimum rating (1-5, default 4)
|
||
|
||
Returns:
|
||
List of dicts with db_id, prompt, rating, feedback, image_path, params.
|
||
"""
|
||
db = _get_db()
|
||
entries = db.search_similar(prompt, limit * 2)
|
||
entries = [e for e in entries if e.rating >= min_rating][:limit]
|
||
|
||
if not entries:
|
||
return []
|
||
|
||
return [
|
||
{
|
||
"db_id": e.id,
|
||
"prompt": e.prompt,
|
||
"rating": e.rating,
|
||
"feedback": e.feedback,
|
||
"image_path": e.image_path,
|
||
"params": e.params,
|
||
}
|
||
for e in entries
|
||
]
|
||
|
||
|
||
@mcp.tool()
|
||
def get_project_style_guide(limit: int = 4) -> dict:
|
||
"""Return the server's current reusable style guide and reference images.
|
||
|
||
Call this before generating a coordinated asset family. The references are
|
||
registered in the MCP server's SQLite DB, rather than hard-coded in a
|
||
client. Pass the returned recommended_reference_images to PixelLab.
|
||
"""
|
||
references = _get_db().get_style_references(limit)
|
||
return {
|
||
"rules": {
|
||
"pixel_art": "Use native 64x64 or 128x128 PNG and integer client scaling.",
|
||
"top_down": "State orthographic camera, nose at 12 o'clock, engines at 6 o'clock.",
|
||
"metadata": "Use one semantic asset role per generation prompt.",
|
||
},
|
||
"recommended_reference_images": [reference.image_path for reference in references],
|
||
"references": [
|
||
{
|
||
"reference_id": reference.id,
|
||
"image_path": reference.image_path,
|
||
"name": reference.name,
|
||
"role": reference.role,
|
||
"notes": reference.notes,
|
||
"priority": reference.priority,
|
||
}
|
||
for reference in references
|
||
],
|
||
}
|
||
|
||
|
||
@mcp.tool()
|
||
def register_style_reference(
|
||
image_path: str,
|
||
name: str,
|
||
role: str,
|
||
notes: Optional[str] = None,
|
||
priority: int = 0,
|
||
) -> dict:
|
||
"""Register an existing generated image as a reusable project style reference.
|
||
|
||
Use a stable, well-reviewed PNG. Higher priority references are returned
|
||
first by get_project_style_guide. Register 1–4 complementary examples, not
|
||
many near-duplicates.
|
||
"""
|
||
if not os.path.isfile(image_path):
|
||
raise ValueError(f"Style reference image does not exist: {image_path}")
|
||
reference_id = _get_db().add_style_reference(
|
||
image_path, name, role, notes, priority
|
||
)
|
||
return {"reference_id": reference_id, "status": "saved", "image_path": image_path}
|
||
|
||
|
||
@mcp.tool()
|
||
def remove_style_reference(reference_id: str) -> dict:
|
||
"""Remove a style-reference registration; the image file is not deleted."""
|
||
_get_db().delete_style_reference(reference_id)
|
||
return {"reference_id": reference_id, "status": "removed"}
|
||
|
||
|
||
@mcp.tool()
|
||
def list_sprites(
|
||
filter: str = "all",
|
||
limit: int = 20,
|
||
) -> list[dict]:
|
||
"""List sprites in the feedback database.
|
||
|
||
Args:
|
||
filter: "all" = all sprites, "unrated" = only unrated, "top" = highest rated
|
||
limit: Max number of results (default 20)
|
||
|
||
Returns:
|
||
List of dicts with db_id, prompt, rating, image_path, created_at.
|
||
"""
|
||
db = _get_db()
|
||
|
||
if filter == "unrated":
|
||
entries = db.get_unrated()
|
||
elif filter == "top":
|
||
entries = db.top_rated(limit, 1)
|
||
else:
|
||
entries = db.get_all()
|
||
|
||
entries = entries[:limit]
|
||
|
||
return [
|
||
{
|
||
"db_id": e.id,
|
||
"prompt": e.prompt,
|
||
"rating": e.rating,
|
||
"image_path": e.image_path,
|
||
"created_at": e.created_at,
|
||
}
|
||
for e in entries
|
||
]
|
||
|
||
|
||
@mcp.tool()
|
||
def db_stats() -> dict:
|
||
"""Get feedback database statistics.
|
||
|
||
Returns:
|
||
Dict with total, rated, unrated, avg_rating.
|
||
"""
|
||
db = _get_db()
|
||
stats = db.stats()
|
||
|
||
return {
|
||
"total": stats.total,
|
||
"rated": stats.rated,
|
||
"unrated": stats.unrated,
|
||
"avg_rating": stats.avg_rating,
|
||
}
|
||
|
||
|
||
if __name__ == "__main__":
|
||
mcp.run(transport="stdio")
|