#!/usr/bin/env python3
"""Run an exported ComfyUI API graph. Python 3.10+, standard library only."""

import argparse
import copy
import json
import math
import sys
import time
import uuid
from pathlib import Path
from urllib.error import HTTPError, URLError
from urllib.parse import quote, urlencode, urlsplit
from urllib.request import Request, urlopen


class WorkflowError(RuntimeError):
    pass


def request(base_url, path, payload=None, timeout=30, binary=False):
    data = None if payload is None else json.dumps(payload).encode('utf-8')
    req = Request(base_url.rstrip('/') + path, data=data,
                  headers={'Content-Type': 'application/json'} if data is not None else {})
    try:
        with urlopen(req, timeout=timeout) as response:
            body = response.read()
    except HTTPError as exc:
        detail = exc.read(4096).decode('utf-8', errors='replace')
        raise WorkflowError(f'HTTP {exc.code} on {path}: {detail}') from exc
    except (URLError, TimeoutError, OSError) as exc:
        raise WorkflowError(f'Request failed on {path}: {exc}') from exc
    if binary:
        return body
    try:
        return json.loads(body)
    except (ValueError, UnicodeDecodeError) as exc:
        raise WorkflowError(f'Expected JSON on {path}; check the URL and access method.') from exc


def load_graph(path):
    graph = json.loads(path.read_text(encoding='utf-8'))
    if not isinstance(graph, dict) or not graph or 'nodes' in graph or 'prompt' in graph:
        raise WorkflowError('Use the .api.json graph, not the canvas JSON or a prompt wrapper.')
    for ident, node in graph.items():
        if not isinstance(node, dict) or not isinstance(node.get('class_type'), str) or not isinstance(node.get('inputs'), dict):
            raise WorkflowError(f'Invalid API node: {ident}')
    return graph


def preflight(graph, info):
    """Check node availability, required inputs, model names and linked slots."""
    if not isinstance(info, dict):
        raise WorkflowError('Invalid /object_info response.')
    for ident, node in graph.items():
        kind = node['class_type']
        if kind not in info:
            raise WorkflowError(f'Node {ident}: {kind} is unavailable on this server.')
        contract = info[kind]
        inputs = node['inputs']
        declared = contract.get('input', {})
        required = declared.get('required', {})
        missing = set(required) - set(inputs)
        if missing:
            raise WorkflowError(f'Node {ident} is missing inputs: {sorted(missing)}')
        specs = {**required, **declared.get('optional', {})}
        for name, value in inputs.items():
            if name not in specs:
                raise WorkflowError(f'Node {ident}: unknown input {name}.')
            expected = specs[name][0]
            if isinstance(value, list):
                if len(value) != 2 or str(value[0]) not in graph or type(value[1]) is not int:
                    raise WorkflowError(f'Node {ident}: invalid link at {name}.')
                source_kind = graph[str(value[0])]['class_type']
                outputs = info.get(source_kind, {}).get('output', [])
                if value[1] < 0 or value[1] >= len(outputs) or outputs[value[1]] != expected:
                    raise WorkflowError(f'Node {ident}: incompatible output slot at {name}.')
            elif isinstance(expected, list) and name != 'image' and value not in expected:
                raise WorkflowError(f'Node {ident}: {name}={value!r} is unavailable; check the model file or dropdown.')
    if not any(n['class_type'] == 'SaveImage' for n in graph.values()):
        raise WorkflowError('This image example requires a SaveImage output node.')


def variation(graph, index):
    result = copy.deepcopy(graph)
    for node in result.values():
        if node['class_type'] in ('KSampler', 'KSamplerAdvanced'):
            key = 'seed' if node['class_type'] == 'KSampler' else 'noise_seed'
            seed = node['inputs'].get(key)
            if type(seed) is not int or not 0 <= seed < 2**64:
                raise WorkflowError(f'{key} must be an integer in the unsigned 64-bit range.')
            node['inputs'][key] = (seed + index) % 2**64
    return result


def run_one(base_url, graph, output_dir, timeout):
    # One POST only. Retrying after a lost response could enqueue a duplicate.
    submitted = request(base_url, '/prompt', {'prompt': graph, 'client_id': str(uuid.uuid4())})
    if not isinstance(submitted, dict) or not submitted.get('prompt_id'):
        raise WorkflowError(f'Graph rejected: {submitted}')
    prompt_id = str(submitted['prompt_id'])
    print(f'Queued prompt_id={prompt_id}', flush=True)
    if submitted.get('error') or submitted.get('node_errors'):
        raise WorkflowError(f'Prompt {prompt_id} was queued with node errors: {submitted}. Inspect it before retrying.')
    deadline = time.monotonic() + timeout
    while True:
        remaining = deadline - time.monotonic()
        if remaining <= 0:
            raise WorkflowError(f'Timed out waiting for {prompt_id}. It may still be queued or running; inspect ComfyUI before retrying.')
        history = request(base_url, '/history/' + quote(prompt_id, safe=''), timeout=min(30, remaining))
        if not isinstance(history, dict):
            raise WorkflowError(f'Invalid history response for {prompt_id}.')
        item = history.get(prompt_id)
        if item is not None:
            status = item.get('status', {})
            messages = status.get('messages', [])
            if status.get('status_str') == 'error' or any(m[0] in ('execution_error', 'execution_interrupted') for m in messages):
                raise WorkflowError(f'Execution failed for {prompt_id}: {messages}')
            if status.get('completed') is True:
                break
        time.sleep(min(1, max(0, deadline - time.monotonic())))

    images = [image for ident, output in item.get('outputs', {}).items()
              if graph.get(str(ident), {}).get('class_type') == 'SaveImage'
              for image in output.get('images', []) if image.get('type') == 'output']
    if not images:
        raise WorkflowError(f'Prompt {prompt_id} completed without saved images.')
    # Generated directory name and local filenames do not trust server paths.
    destination = output_dir / uuid.uuid4().hex
    destination.mkdir(parents=True, exist_ok=False)
    (destination / 'request.json').write_text(json.dumps({'prompt_id': prompt_id, 'prompt': graph}, indent=2) + '\n', encoding='utf-8')
    for index, image in enumerate(images):
        query = urlencode({key: image.get(key, '') for key in ('filename', 'subfolder', 'type')})
        data = request(base_url, '/view?' + query, binary=True)
        if not data.startswith(b'\x89PNG\r\n\x1a\n'):
            raise WorkflowError(f'Expected a PNG from SaveImage for {prompt_id}; download was not an image.')
        (destination / f'{index + 1:03d}.png').write_bytes(data)
    print(f'Saved {len(images)} image(s) to {destination}', flush=True)
    return destination


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('workflow', type=Path, help='Exported API graph (.api.json)')
    parser.add_argument('--base-url', default='http://127.0.0.1:8188')
    parser.add_argument('--output-dir', type=Path, default=Path('comfy-output'))
    parser.add_argument('--timeout', type=float, default=600, help='Seconds to wait per queued job')
    parser.add_argument('--count', type=int, default=1, help='Sequential runs, incrementing sampler seeds (1–100)')
    parser.add_argument('--check', action='store_true', help='Read /object_info only; do not queue GPU work')
    args = parser.parse_args()
    url = urlsplit(args.base_url)
    if url.scheme not in ('http', 'https') or not url.netloc or url.username or url.password or url.query or url.fragment:
        parser.error('Use an HTTP(S) server URL without embedded credentials, query or fragment.')
    if not math.isfinite(args.timeout) or args.timeout <= 0 or not 1 <= args.count <= 100:
        parser.error('--timeout must be finite and positive; --count must be 1–100.')
    try:
        graph = load_graph(args.workflow)
        preflight(graph, request(args.base_url, '/object_info'))
        if args.check:
            print('Node contracts and model dropdowns match. This does not test GPU memory, execution or image quality.')
            return 0
        if args.count > 1 and not any(n['class_type'] in ('KSampler', 'KSamplerAdvanced') for n in graph.values()):
            raise WorkflowError('This graph has no sampler to vary; run it once per input image.')
        for index in range(args.count):
            run_one(args.base_url, variation(graph, index), args.output_dir, args.timeout)
        return 0
    except (WorkflowError, OSError, ValueError) as exc:
        print(f'Error: {exc}', file=sys.stderr)
        return 1


if __name__ == '__main__':
    sys.exit(main())
