import json
import os
import sys
import time
import urllib.error
import urllib.parse
import urllib.request

API_BASE = os.environ.get("UDIOAPI_BASE_URL", "https://udioapi.pro/api").rstrip("/")
API_KEY = os.environ.get("UDIOAPI_KEY")
POLL_INTERVAL_SECONDS = 5
MAX_POLLS = 60

if not API_KEY:
    raise RuntimeError("Set UDIOAPI_KEY in the server environment before running this example.")


def request_json(path, method="GET", body=None):
    encoded_body = json.dumps(body).encode("utf-8") if body is not None else None
    request = urllib.request.Request(
        API_BASE + path,
        data=encoded_body,
        method=method,
        headers={
            "Authorization": "Bearer " + API_KEY,
            "Content-Type": "application/json",
        },
    )

    try:
        with urllib.request.urlopen(request, timeout=30) as response:
            status = response.status
            raw = response.read().decode("utf-8")
    except urllib.error.HTTPError as error:
        status = error.code
        raw = error.read().decode("utf-8")

    try:
        payload = json.loads(raw) if raw else {}
    except json.JSONDecodeError as error:
        raise RuntimeError("The API returned a non-JSON response with HTTP " + str(status) + ".") from error

    if status >= 400 or int(payload.get("code") or 0) >= 400:
        message = payload.get("message") or payload.get("error") or "API request failed"
        raise RuntimeError("HTTP " + str(status) + ": " + str(message))

    return payload


def read_work_id(payload):
    data = payload.get("data") if isinstance(payload.get("data"), dict) else {}
    return (
        payload.get("workId")
        or payload.get("task_id")
        or data.get("workId")
        or data.get("task_id")
    )


def inspect_task(payload):
    data = payload.get("data") if isinstance(payload.get("data"), dict) else payload
    items = data.get("response_data") or data.get("items") or []
    if not isinstance(items, list):
        items = []

    for item in items:
        if not isinstance(item, dict):
            continue
        failure = item.get("fail_message") or item.get("error_message")
        if failure:
            return True, str(failure), []

    status = str(data.get("status") or "").upper()
    if status in {"FAILED", "ERROR"}:
        failure = (
            data.get("fail_message")
            or data.get("error_message")
            or data.get("message")
            or data.get("error")
            or ("Music task ended with status " + status)
        )
        return True, str(failure), []

    audio_urls = list(dict.fromkeys(
        item.get("audio_url")
        for item in items
        if isinstance(item, dict) and item.get("audio_url")
    ))
    terminal_message = any(
        "all generated successfully" in str(item.get("extra_message") or "").lower()
        for item in items
        if isinstance(item, dict)
    )
    terminal_status = status in {
        "SUCCESS",
        "COMPLETED",
        "COMPLETE",
    }
    return terminal_message or terminal_status or len(audio_urls) >= 2, None, audio_urls


def main():
    created = request_json(
        "/v2/generate",
        method="POST",
        body={
            "model": "chirp-v5-5",
            "gpt_description_prompt": "Warm indie pop with acoustic guitar and a hopeful chorus",
            "make_instrumental": True,
        },
    )
    work_id = read_work_id(created)
    if not work_id:
        raise RuntimeError("The create response did not include workId or task_id.")

    print("Created music task:", work_id)

    for attempt in range(1, MAX_POLLS + 1):
        query = urllib.parse.urlencode({"workId": work_id})
        status = request_json("/v2/feed?" + query)
        done, failure, audio_urls = inspect_task(status)

        if failure:
            raise RuntimeError("Music task failed: " + failure)
        if done:
            print("Completed audio URLs:", audio_urls)
            return

        print("Still processing (" + str(attempt) + "/" + str(MAX_POLLS) + ")")
        time.sleep(POLL_INTERVAL_SECONDS)

    raise RuntimeError("Polling timed out. Save the workId and query the feed endpoint again later.")


if __name__ == "__main__":
    try:
        main()
    except Exception as error:
        print(str(error), file=sys.stderr)
        sys.exit(1)
