"""One-request OpenAI Responses stream diagnostic. See article setup and limits."""
import asyncio
import json
import os
import time
from datetime import datetime, timezone

import aiohttp

OPENAI_RESPONSES_URL = "https://api.openai.com/v1/responses"
CONNECT_LIMIT_SECONDS = 5.0
SOCKET_CONNECT_LIMIT_SECONDS = 5.0
SOCKET_READ_IDLE_SECONDS = 30.0
TOTAL_DEADLINE_SECONDS = 90.0


def utc_now():
    return datetime.now(timezone.utc).isoformat(timespec="milliseconds")


def safe_exception_chain(exc):
    """Return class names only; never serialize exception text or headers."""
    names = []
    seen = set()
    current = exc
    while current is not None and id(current) not in seen and len(names) < 5:
        seen.add(id(current))
        names.append(type(current).__name__)
        current = current.__cause__
    return names


def parse_sse_data(data_lines):
    if not data_lines:
        return None
    data = "\n".join(data_lines)
    if data == "[DONE]":
        return {"type": "sse_done_marker"}
    try:
        event = json.loads(data)
    except json.JSONDecodeError:
        return {"type": "unparsed_sse_data"}
    if isinstance(event, dict):
        return event
    return {"type": "non_object_sse_data"}


async def run_probe(url, token, model, *, total_seconds=TOTAL_DEADLINE_SECONDS,
                    sock_read_seconds=SOCKET_READ_IDLE_SECONDS):
    started = time.monotonic()
    state = {
        "status": None,
        "request_id": None,
        "phase": "request_start",
        "first_event_at": None,
        "first_event_seconds": None,
        "last_event_type": None,
        "last_event_at": None,
        "terminal_event": None,
        "application_outcome": None,
        "exception_classes": [],
        "deadline_expired": False,
        "sse_frame_pending": False,
    }

    def log(record):
        print(json.dumps(record, separators=(",", ":")))

    log({"stage": "request_start", "at": utc_now(), "model": model,
         "retry_count": 0, "timeouts_s": {"connect_including_pool": CONNECT_LIMIT_SECONDS,
         "socket_connect": SOCKET_CONNECT_LIMIT_SECONDS,
         "read_between_chunks": sock_read_seconds,
         "total_application_deadline": total_seconds}})

    deadline = None
    cancellation = None
    try:
        timeout = aiohttp.ClientTimeout(
            total=None,
            connect=CONNECT_LIMIT_SECONDS,
            sock_connect=SOCKET_CONNECT_LIMIT_SECONDS,
            sock_read=sock_read_seconds,
        )
        headers = {"Authorization": "Bearer " + token,
                   "Content-Type": "application/json",
                   "Accept": "text/event-stream"}
        payload = {"model": model, "input": "Reply with the single word OK.",
                   "stream": True}
        async with aiohttp.ClientSession(timeout=timeout, trust_env=False) as session:
            deadline = asyncio.timeout(total_seconds)
            try:
                async with deadline:
                    state["phase"] = "waiting_for_response_headers"
                    async with session.post(url, headers=headers, json=payload,
                                            allow_redirects=False) as response:
                        state["status"] = response.status
                        state["request_id"] = response.headers.get("x-request-id")
                        state["phase"] = "response_body"
                        log({"stage": "headers_received", "at": utc_now(),
                             "seconds_from_start": round(time.monotonic() - started, 3),
                             "status": state["status"],
                             "request_id": state["request_id"]})

                        if not 200 <= response.status < 300:
                            state["application_outcome"] = "http_rejection"
                            state["phase"] = "http_rejection_body_not_read"
                        else:
                            data_lines = []
                            async for line_bytes in response.content:
                                state["phase"] = "response_body"
                                line = line_bytes.decode("utf-8", errors="replace").rstrip("\r\n")
                                if line == "":
                                    event = parse_sse_data(data_lines)
                                    data_lines = []
                                    state["sse_frame_pending"] = False
                                    if event is None:
                                        continue
                                    now = utc_now()
                                    event_type = event.get("type", "unknown")
                                    state["last_event_type"] = event_type
                                    state["last_event_at"] = now
                                    if state["first_event_at"] is None:
                                        state["first_event_at"] = now
                                        state["first_event_seconds"] = round(time.monotonic() - started, 3)
                                        log({"stage": "first_event", "at": now,
                                             "seconds_from_start": state["first_event_seconds"],
                                             "event_type": event_type})
                                    if event_type in {"response.completed", "response.incomplete",
                                                      "response.failed", "error"}:
                                        state["terminal_event"] = event_type
                                        if event_type == "response.completed":
                                            state["application_outcome"] = "complete"
                                        elif event_type == "response.incomplete":
                                            state["application_outcome"] = "incomplete"
                                        else:
                                            state["application_outcome"] = "application_error"
                                    log({"stage": "event_observed", "at": now,
                                         "event_type": event_type,
                                         "request_id": state["request_id"]})
                                elif line.startswith("data:"):
                                    value = line[5:]
                                    data_lines.append(value[1:] if value.startswith(" ") else value)
                                    state["sse_frame_pending"] = True
            except TimeoutError as exc:
                state["deadline_expired"] = bool(deadline and deadline.expired())
                state["exception_classes"] = safe_exception_chain(exc)
                if state["deadline_expired"]:
                    state["phase"] = "total_application_deadline"
                else:
                    state["phase"] = "transport_timeout"
                if state["application_outcome"] is None:
                    state["application_outcome"] = "unknown_after_request_start"
            except Exception as exc:
                state["exception_classes"] = safe_exception_chain(exc)
                state["phase"] = "transport_or_stream_error"
                if state["application_outcome"] is None:
                    state["application_outcome"] = "unknown_after_request_start"
    except asyncio.CancelledError as exc:
        cancellation = exc
        state["exception_classes"] = safe_exception_chain(exc)
        state["phase"] = "caller_cancelled"
        if state["application_outcome"] is None:
            state["application_outcome"] = "unknown_after_request_start"
    except Exception as exc:
        state["exception_classes"] = safe_exception_chain(exc)
        state["phase"] = "client_setup_or_request_error"
        if state["application_outcome"] is None:
            state["application_outcome"] = "unknown_after_request_start"

    if state["application_outcome"] is None:
        state["application_outcome"] = (
            "stream_ended_with_unterminated_sse_frame" if state["sse_frame_pending"]
            else "stream_ended_without_terminal_event")
    log({"stage": "request_end", "at": utc_now(),
         "phase": state["phase"], "application_outcome": state["application_outcome"],
         "status": state["status"], "request_id": state["request_id"],
         "seconds_total": round(time.monotonic() - started, 3),
         "seconds_to_first_event": state["first_event_seconds"],
         "last_event_type": state["last_event_type"],
         "last_event_at": state["last_event_at"],
         "terminal_event": state["terminal_event"],
         "exception_classes": state["exception_classes"],
         "total_deadline_expired": state["deadline_expired"],
         "unterminated_sse_frame": state["sse_frame_pending"]})
    if cancellation is not None:
        raise cancellation
    return state


async def main():
    token = os.environ["OPENAI_API_KEY"]
    model = os.environ["OPENAI_MODEL"]
    await run_probe(OPENAI_RESPONSES_URL, token, model)


if __name__ == "__main__":
    asyncio.run(main())
