#!/usr/bin/env python3 """Prompt enhancement (PE) for Ming-Image text-to-image via a Ling-3.0-flash-VL seat. Per the README, PE is a pre-processing step *outside* ``infer.py``: an instruction-following VLM rewrites a short caption into the structured Figma-style JSON prompt that the text-to-image pipeline consumes, and the result is passed to ``infer.py --prompt`` as raw text or via a file. This module drives any OpenAI-compatible ``/chat/completions`` endpoint using only the standard library (``urllib``): by default the local llama-server seat serving Ling-3.0-flash-VL, optionally the LiteLLM lab gateway (Bearer auth via ``--api-key`` or the ``LITELLM_API_KEY`` environment variable). The rewriter system prompt is read verbatim from ``assets/t2i_rewriter_system_prompt.txt``. The reply is parsed robustly (```json fences and surrounding prose are tolerated), then validated against the schema the system prompt demands. On a parse or validation failure the request is retried exactly once with the errors appended to the user turn; if that still fails, PromptEnhancementError is raised with the errors. Invalid JSON is never passed through silently. CLI: python pe_ling.py "a caption" --out prompt.json \ [--base-url http://127.0.0.1:8090/v1] \ [--model ling-3.0-flash-vl-mtp-halo-STRIX_LEAN] """ from __future__ import annotations import argparse import json import os import re import sys import time import urllib.error import urllib.request from pathlib import Path from typing import Any, Dict, List, Optional, Tuple CODE_DIRECTORY = Path(__file__).resolve().parent SYSTEM_PROMPT_PATH = CODE_DIRECTORY / "assets" / "t2i_rewriter_system_prompt.txt" # The Ling-3.0-flash-VL seat already served on the target box (llama-server, # OpenAI-compatible, thinking disabled); both endpoints speak the same # /chat/completions protocol. DEFAULT_BASE_URL = "http://127.0.0.1:8090/v1" DEFAULT_MODEL = "ling-3.0-flash-vl-mtp-halo-STRIX_LEAN" API_KEY_ENV = "LITELLM_API_KEY" # Low temperature: the rewrite is a deterministic schema transformation, not # creative sampling. DEFAULT_TEMPERATURE = 0.2 # The upstream example rewrite (assets/t2i_four_seasons_cabin_prompt.json) is # ~5 KB (~2k tokens); dense multi-layer infographic rewrites run several times # longer, so leave generous headroom for a complete JSON object. DEFAULT_MAX_TOKENS = 16384 # A multi-thousand-token completion on the local seat can take minutes. DEFAULT_TIMEOUT_SECONDS = 600.0 REPAIR_INSTRUCTION = "Return only the corrected JSON object: no prose, no code fences." CANVAS_SETTINGS_KEYS = ("aspect_ratio", "ambient_lighting", "image_style") LAYER_KEYS = ("description", "coordinates", "hierarchy_and_relation", "color_specs") COORDINATE_FIELDS = ("cx", "cy", "w", "h") # `coordinates` must be ONE string of the form # "cx: 0.500, cy: 0.500, w: 1.000, h: 1.000". The upstream example also uses # bare integers ("h: 1"), so accept any decimal spelling and enforce the # [0, 1] range on the parsed value. Whitespace around ':' and ',' is # tolerated; the key order is fixed. _COORDINATE_NUMBER = r"[-+]?(?:\d+(?:\.\d*)?|\.\d+)" COORDINATES_RE = re.compile( rf"^\s*cx:\s*(?P{_COORDINATE_NUMBER})\s*," rf"\s*cy:\s*(?P{_COORDINATE_NUMBER})\s*," rf"\s*w:\s*(?P{_COORDINATE_NUMBER})\s*," rf"\s*h:\s*(?P{_COORDINATE_NUMBER})\s*$" ) # Hex colors: #RGB, #RGBA, #RRGGBB, #RRGGBBAA (the upstream example uses # #RRGGBB; the alpha forms keep RGBA-design outputs from failing validation). HEX_COLOR_RE = re.compile( r"^#(?:[0-9a-fA-F]{3}|[0-9a-fA-F]{4}|[0-9a-fA-F]{6}|[0-9a-fA-F]{8})$" ) class PromptEnhancementError(RuntimeError): """PE failed: transport/protocol error, or schema failure after the retry.""" def __init__( self, message: str, errors: Optional[List[str]] = None, reply: Optional[str] = None, ): super().__init__(message) self.errors = list(errors or []) self.reply = reply def load_system_prompt(path: Path = SYSTEM_PROMPT_PATH) -> str: """Return the released rewriter system prompt, verbatim.""" return path.read_text(encoding="utf-8") def extract_json_object(text: str) -> Dict[str, Any]: """Return the first complete top-level JSON object found in ``text``. Models sometimes wrap JSON in ```json fences or add prose around it. Scanning every ``{`` position with ``JSONDecoder.raw_decode`` (which decodes a document at an offset and ignores trailing data) recovers the object in all of those shapes. Raises ValueError when no complete JSON object is present, e.g. a reply truncated mid-object. """ decoder = json.JSONDecoder() position = text.find("{") while position != -1: try: document, _ = decoder.raw_decode(text, position) except ValueError: position = text.find("{", position + 1) continue return document snippet = text.strip() if len(snippet) > 300: snippet = snippet[:300] + "..." raise ValueError( f"reply contains no complete top-level JSON object " f"({len(text)} characters); starts with: {snippet!r}" ) def _check_exact_keys( mapping: Dict[str, Any], expected: Tuple[str, ...], path: str, errors: List[str] ) -> None: missing = [key for key in expected if key not in mapping] unexpected = [key for key in mapping if key not in expected] if missing: errors.append(f"{path}: missing required key(s): {', '.join(missing)}") if unexpected: errors.append( f"{path}: unexpected key(s): {', '.join(unexpected)} " f"(exactly {', '.join(expected)} are required)" ) def _check_non_empty_string(value: Any, path: str, errors: List[str]) -> None: if not isinstance(value, str): errors.append(f"{path}: expected a string, got {type(value).__name__}") elif not value.strip(): errors.append(f"{path}: string is empty") def _check_coordinates(value: Any, path: str, errors: List[str]) -> None: if not isinstance(value, str): errors.append( f"{path}: must be ONE string of the form " f"'cx: 0.500, cy: 0.500, w: 1.000, h: 1.000', got {type(value).__name__}" ) return match = COORDINATES_RE.match(value) if match is None: errors.append( f"{path}: {value!r} is not of the form " f"'cx: 0.500, cy: 0.500, w: 1.000, h: 1.000'" ) return for field in COORDINATE_FIELDS: number = float(match.group(field)) if not 0.0 <= number <= 1.0: errors.append(f"{path}: {field}={match.group(field)} is outside [0, 1]") def _check_color_specs(value: Any, path: str, errors: List[str]) -> None: if not isinstance(value, list): errors.append( f"{path}: expected a list of hex colors, got {type(value).__name__}" ) return for index, color in enumerate(value): if not isinstance(color, str) or HEX_COLOR_RE.match(color) is None: errors.append( f"{path}[{index}]: {color!r} is not a hex color " f"(expected #RGB, #RGBA, #RRGGBB, or #RRGGBBAA)" ) def validate_enhanced_prompt(document: Any) -> List[str]: """Return schema errors for a rewritten prompt; an empty list means valid. Schema demanded by assets/t2i_rewriter_system_prompt.txt: exactly two top-level keys ``canvas_settings`` (exactly ``aspect_ratio``, ``ambient_lighting``, ``image_style``) and ``layers`` (each layer exactly ``description``, ``coordinates``, ``hierarchy_and_relation``, ``color_specs``); ``coordinates`` is a string "cx: 0.500, cy: 0.500, w: 1.000, h: 1.000" with values in [0, 1]; ``color_specs`` is a list of hex colors. ``layers`` must hold at least one visible layer -- an empty list means the rewrite failed even though it is type-correct. """ if not isinstance(document, dict): return [f"top level: expected a JSON object, got {type(document).__name__}"] errors: List[str] = [] _check_exact_keys(document, ("canvas_settings", "layers"), "top level", errors) if "canvas_settings" in document: canvas = document["canvas_settings"] if not isinstance(canvas, dict): errors.append( f"canvas_settings: expected a JSON object, got {type(canvas).__name__}" ) else: _check_exact_keys(canvas, CANVAS_SETTINGS_KEYS, "canvas_settings", errors) for key in CANVAS_SETTINGS_KEYS: if key in canvas: _check_non_empty_string( canvas[key], f"canvas_settings.{key}", errors ) if "layers" in document: layers = document["layers"] if not isinstance(layers, list): errors.append(f"layers: expected a list, got {type(layers).__name__}") elif not layers: errors.append("layers: expected at least one visible layer") else: for index, layer in enumerate(layers): path = f"layers[{index}]" if not isinstance(layer, dict): errors.append( f"{path}: expected a JSON object, got {type(layer).__name__}" ) continue _check_exact_keys(layer, LAYER_KEYS, path, errors) for key in ("description", "hierarchy_and_relation"): if key in layer: _check_non_empty_string(layer[key], f"{path}.{key}", errors) if "coordinates" in layer: _check_coordinates( layer["coordinates"], f"{path}.coordinates", errors ) if "color_specs" in layer: _check_color_specs(layer["color_specs"], f"{path}.color_specs", errors) return errors def _chat_completion( base_url: str, model: str, messages: List[Dict[str, str]], *, temperature: float, max_tokens: int, api_key: Optional[str], timeout: float, ) -> Tuple[str, Optional[str]]: """POST one chat completion; return (content, finish_reason).""" url = base_url.rstrip("/") + "/chat/completions" payload = json.dumps( { "model": model, "messages": messages, "temperature": temperature, "max_tokens": max_tokens, "stream": False, } ).encode("utf-8") headers = {"Content-Type": "application/json"} if api_key: headers["Authorization"] = f"Bearer {api_key}" request = urllib.request.Request(url, data=payload, headers=headers, method="POST") try: with urllib.request.urlopen(request, timeout=timeout) as response: body = response.read().decode("utf-8", errors="replace") except urllib.error.HTTPError as error: detail = error.read().decode("utf-8", errors="replace") raise PromptEnhancementError( f"HTTP {error.code} from {url}: {detail[:2000]}" ) from error except urllib.error.URLError as error: raise PromptEnhancementError(f"cannot reach {url}: {error.reason}") from error except OSError as error: # includes socket timeouts during the read raise PromptEnhancementError(f"request to {url} failed: {error}") from error try: envelope = json.loads(body) choice = envelope["choices"][0] content = choice["message"]["content"] except (json.JSONDecodeError, KeyError, IndexError, TypeError) as error: raise PromptEnhancementError( f"malformed chat completion response from {url}: {body[:500]}" ) from error finish_reason = choice.get("finish_reason") if not isinstance(content, str) or not content.strip(): raise PromptEnhancementError( f"empty completion content from {url} (finish_reason={finish_reason!r})" ) return content, finish_reason def enhance( caption: str, base_url: str, model: str, api_key: Optional[str] = None, timeout: float = DEFAULT_TIMEOUT_SECONDS, temperature: float = DEFAULT_TEMPERATURE, max_tokens: int = DEFAULT_MAX_TOKENS, ) -> Dict[str, Any]: """Return the validated structured rewrite of ``caption``. Sends the verbatim rewriter system prompt plus the caption to ``{base_url}/chat/completions``. On a parse or schema failure, retries exactly once with the validation errors appended to the user turn; if that also fails, raises PromptEnhancementError carrying the errors. """ system_prompt = load_system_prompt() messages = [ {"role": "system", "content": system_prompt}, {"role": "user", "content": caption}, ] request_kwargs = { "temperature": temperature, "max_tokens": max_tokens, "api_key": api_key, "timeout": timeout, } errors: List[str] = [] content = "" for attempt in (1, 2): content, finish_reason = _chat_completion( base_url, model, messages, **request_kwargs ) document: Optional[Dict[str, Any]] = None try: document = extract_json_object(content) except ValueError as error: errors = [str(error)] if document is not None: errors = validate_enhanced_prompt(document) if not errors: assert document is not None # errors empty implies extraction succeeded return document if finish_reason == "length": errors.append( "the reply was cut off (finish_reason='length'): the complete " f"JSON object must fit within max_tokens={max_tokens}" ) print(f"pe_ling: attempt {attempt}/2 failed validation:", file=sys.stderr) for error in errors: print(f"pe_ling: - {error}", file=sys.stderr) if attempt == 1: retry_content = ( f"{caption}\n\n" "Your previous reply failed schema validation:\n" + "".join(f"- {error}\n" for error in errors) + "\n" + REPAIR_INSTRUCTION ) messages = [ {"role": "system", "content": system_prompt}, {"role": "user", "content": retry_content}, ] raise PromptEnhancementError( "prompt enhancement failed schema validation after 2 attempts:\n" + "".join(f" - {error}\n" for error in errors).rstrip(), errors=errors, reply=content, ) def main() -> None: parser = argparse.ArgumentParser( description=( "Enhance a Ming-Image text-to-image caption into the validated " "structured JSON prompt via an OpenAI-compatible Ling-3.0-flash-VL " "endpoint." ) ) parser.add_argument("caption", help="free-form design caption to enhance") parser.add_argument( "--out", type=Path, help="write the validated JSON here (default: stdout, summary on stderr)", ) parser.add_argument( "--base-url", default=DEFAULT_BASE_URL, help=f"OpenAI-compatible base URL (default: {DEFAULT_BASE_URL})", ) parser.add_argument( "--model", default=DEFAULT_MODEL, help=f"chat model id served at the endpoint (default: {DEFAULT_MODEL})", ) parser.add_argument( "--api-key", default=os.environ.get(API_KEY_ENV), help=f"Bearer token for gated endpoints; defaults to ${API_KEY_ENV} when set", ) parser.add_argument( "--timeout", type=float, default=DEFAULT_TIMEOUT_SECONDS, help=f"per-request timeout in seconds (default: {DEFAULT_TIMEOUT_SECONDS})", ) parser.add_argument( "--temperature", type=float, default=DEFAULT_TEMPERATURE, help=f"sampling temperature (default: {DEFAULT_TEMPERATURE})", ) parser.add_argument( "--max-tokens", type=int, default=DEFAULT_MAX_TOKENS, help=f"completion token budget (default: {DEFAULT_MAX_TOKENS})", ) args = parser.parse_args() started = time.perf_counter() try: document = enhance( args.caption, args.base_url, args.model, api_key=args.api_key, timeout=args.timeout, temperature=args.temperature, max_tokens=args.max_tokens, ) except PromptEnhancementError as error: print(f"pe_ling: {error}", file=sys.stderr) raise SystemExit(1) elapsed = time.perf_counter() - started layer_count = len(document["layers"]) payload = json.dumps(document, indent=2, ensure_ascii=False) + "\n" if args.out is not None: args.out.parent.mkdir(parents=True, exist_ok=True) args.out.write_text(payload, encoding="utf-8") print(f"pe_ling: {elapsed:.1f}s, {layer_count} layer(s) -> {args.out}") else: sys.stdout.write(payload) print(f"pe_ling: {elapsed:.1f}s, {layer_count} layer(s)", file=sys.stderr) if __name__ == "__main__": main()