#!/usr/bin/env python3 """Submit a Chorus text-to-image task, poll it, and download its first image.""" from __future__ import annotations import argparse import ipaddress import json import os import sys import time import uuid from pathlib import Path from urllib.error import HTTPError, URLError from urllib.parse import urljoin, urlsplit from urllib.request import HTTPRedirectHandler, ProxyHandler, Request, build_opener class ChorusError(RuntimeError): pass def origin(url: str) -> tuple[str, str, int | None]: parsed = urlsplit(url) port = parsed.port if port is None: port = 443 if parsed.scheme.lower() == "https" else 80 if parsed.scheme.lower() == "http" else None return parsed.scheme.lower(), (parsed.hostname or "").lower(), port class SameOriginRedirectHandler(HTTPRedirectHandler): def redirect_request(self, request, file_pointer, code, message, headers, new_url): if origin(request.full_url) != origin(new_url): raise ChorusError("Chorus attempted a cross-origin redirect") return super().redirect_request(request, file_pointer, code, message, headers, new_url) def open_url(request: Request, timeout: float): host = urlsplit(request.full_url).hostname or "" try: loopback = ipaddress.ip_address(host).is_loopback except ValueError: loopback = host.lower() == "localhost" if loopback: return build_opener(ProxyHandler({}), SameOriginRedirectHandler()).open(request, timeout=timeout) return build_opener(SameOriginRedirectHandler()).open(request, timeout=timeout) def read_env(path: Path) -> dict[str, str]: if not path.is_file(): raise ChorusError(f"environment file not found: {path}") values: dict[str, str] = {} for number, raw_line in enumerate(path.read_text(encoding="utf-8-sig").splitlines(), 1): line = raw_line.strip() if not line or line.startswith("#"): continue if line.lower().startswith("export "): line = line[7:].lstrip() if "=" not in line: raise ChorusError(f"invalid environment entry at line {number}") key, value = line.split("=", 1) key = key.strip().upper() value = value.strip() if len(value) >= 2 and value[0] == value[-1] and value[0] in "\"'": value = value[1:-1] values[key] = value return values def required_setting(values: dict[str, str], names: tuple[str, ...], label: str) -> str: for name in names: value = values.get(name, "").strip() if value: return value raise ChorusError(f"{label} is missing from chorus.env") def output_path(args: argparse.Namespace) -> Path: if args.output: if args.output_dir or args.filename: raise ChorusError("use --output or --output-dir with --filename, not both") return Path(args.output).expanduser().resolve() if not args.output_dir or not args.filename: raise ChorusError("provide --output or both --output-dir and --filename") if Path(args.filename).name != args.filename: raise ChorusError("--filename must not contain directory components") return (Path(args.output_dir).expanduser() / args.filename).resolve() def multipart_body(prompt: str) -> tuple[bytes, str]: boundary = f"chorus-{uuid.uuid4().hex}" metadata = json.dumps( {"capability": "image_generate", "role_rule": "", "images": []}, ensure_ascii=False, separators=(",", ":"), ) parts: list[bytes] = [] for name, value, content_type in ( ("prompt", prompt, "text/plain; charset=utf-8"), ("metadata", metadata, "application/json; charset=utf-8"), ): parts.extend( [ f"--{boundary}\r\n".encode(), f'Content-Disposition: form-data; name="{name}"\r\n'.encode(), f"Content-Type: {content_type}\r\n\r\n".encode(), value.encode("utf-8"), b"\r\n", ] ) parts.append(f"--{boundary}--\r\n".encode()) return b"".join(parts), boundary def error_message(payload: bytes, fallback: str) -> str: try: parsed = json.loads(payload.decode("utf-8")) error = parsed.get("error", {}) if isinstance(parsed, dict) else {} code = error.get("code") message = error.get("message") if code and message: return f"{code}: {message}" if message: return str(message) except (UnicodeDecodeError, json.JSONDecodeError, AttributeError): pass return fallback def request_json( url: str, api_key: str, method: str = "GET", body: bytes | None = None, headers: dict[str, str] | None = None, timeout: float = 30, ) -> dict: request_headers = {"Accept": "application/json", "Authorization": f"Bearer {api_key}"} if headers: request_headers.update(headers) request = Request(url, data=body, headers=request_headers, method=method) try: with open_url(request, timeout=timeout) as response: payload = response.read() except HTTPError as exc: payload = exc.read() raise ChorusError(f"Chorus HTTP {exc.code}: {error_message(payload, exc.reason)}") from None except URLError as exc: raise ChorusError(f"cannot reach Chorus: {exc.reason}") from None try: parsed = json.loads(payload.decode("utf-8")) except (UnicodeDecodeError, json.JSONDecodeError) as exc: raise ChorusError("Chorus returned invalid JSON") from exc if not isinstance(parsed, dict): raise ChorusError("Chorus returned an unexpected response") return parsed def submit(base_url: str, api_key: str, prompt: str, idempotency_key: str) -> dict: body, boundary = multipart_body(prompt) return request_json( f"{base_url}/openapi/v1/generations/image", api_key, method="POST", body=body, headers={ "Content-Type": f"multipart/form-data; boundary={boundary}", "Idempotency-Key": idempotency_key, }, ) def wait_for_result( base_url: str, api_key: str, generation_id: int, timeout: float, poll_interval: float, ) -> dict: deadline = time.monotonic() + timeout while True: generation = request_json( f"{base_url}/openapi/v1/generations/{generation_id}", api_key ) status = generation.get("status") if status == "succeeded": return generation if status == "failed": code = generation.get("error_code") or "generation_failed" message = generation.get("error_message") or "generation failed" raise ChorusError(f"{code}: {message}") if status not in {"pending", "running"}: raise ChorusError(f"unexpected generation status: {status!r}") if time.monotonic() >= deadline: raise ChorusError(f"generation {generation_id} did not finish within {timeout:g} seconds") time.sleep(min(poll_interval, max(0.0, deadline - time.monotonic()))) def choose_image(generation: dict) -> str: outputs = generation.get("outputs") if not isinstance(outputs, list): raise ChorusError("completed generation has no outputs") for output in outputs: if isinstance(output, dict) and output.get("kind") == "image" and output.get("url"): return str(output["url"]) raise ChorusError("completed generation has no image output") def download(url: str, api_key: str, destination: Path, overwrite: bool) -> None: if destination.exists() and not overwrite: raise ChorusError(f"output already exists: {destination}") destination.parent.mkdir(parents=True, exist_ok=True) temporary = destination.with_name(f".{destination.name}.{uuid.uuid4().hex}.part") request = Request(url, headers={"Accept": "image/*", "Authorization": f"Bearer {api_key}"}) try: with open_url(request, timeout=60) as response, temporary.open("wb") as output: content_type = response.headers.get_content_type() if not content_type.startswith("image/"): raise ChorusError(f"output is not an image: {content_type}") while chunk := response.read(1024 * 1024): output.write(chunk) os.replace(temporary, destination) except HTTPError as exc: raise ChorusError(f"image download HTTP {exc.code}: {error_message(exc.read(), exc.reason)}") from None except URLError as exc: raise ChorusError(f"cannot download image: {exc.reason}") from None finally: if temporary.exists(): temporary.unlink() def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--env", default="chorus.env", help="path to chorus.env") parser.add_argument("--prompt", required=True) parser.add_argument("--output") parser.add_argument("--output-dir") parser.add_argument("--filename") parser.add_argument("--timeout", type=float) parser.add_argument("--poll-interval", type=float) parser.add_argument("--idempotency-key") parser.add_argument("--overwrite", action="store_true") return parser.parse_args() def main() -> int: args = parse_args() try: values = read_env(Path(args.env).expanduser().resolve()) base_url = required_setting( values, ("CHORUS_URL", "CHORUS_API_URL", "URL"), "Chorus URL" ).rstrip("/") api_key = required_setting( values, ("CHORUS_API_KEY", "CHORUS_APIKEY", "APIKEY", "API_KEY"), "Chorus API key", ) destination = output_path(args) if destination.exists() and not args.overwrite: raise ChorusError(f"output already exists: {destination}") timeout = args.timeout if args.timeout is not None else float( values.get("CHORUS_TIMEOUT_SECONDS", "600") ) poll_interval = args.poll_interval if args.poll_interval is not None else float( values.get("CHORUS_POLL_INTERVAL_SECONDS", "2") ) if not args.prompt.strip(): raise ChorusError("prompt must not be empty") if timeout <= 0 or poll_interval <= 0: raise ChorusError("timeout and poll interval must be positive") idempotency_key = args.idempotency_key or f"chorus-image-{uuid.uuid4().hex}" generation = submit(base_url, api_key, args.prompt, idempotency_key) generation_id = generation.get("id") if not isinstance(generation_id, int) or generation_id <= 0: raise ChorusError("submission response has no valid generation ID") print(f"Chorus generation {generation_id} submitted", flush=True) generation = wait_for_result(base_url, api_key, generation_id, timeout, poll_interval) image_url = urljoin(f"{base_url}/", choose_image(generation)) if origin(image_url) != origin(base_url): raise ChorusError("Chorus returned a cross-origin image URL") download(image_url, api_key, destination, args.overwrite) print(f"Saved Chorus generation {generation_id} to {destination}") return 0 except (ChorusError, OSError, ValueError) as exc: print(f"chorus: {exc}", file=sys.stderr) return 1 if __name__ == "__main__": raise SystemExit(main())