from __future__ import annotations

import json
import re

import pytest


def _parse_sse_events(text: str) -> list[tuple[str, dict]]:
    events: list[tuple[str, dict]] = []
    for block in re.split(r"\n\n+", text.strip()):
        if not block:
            continue
        event_type = "message"
        data_line = None
        for line in block.splitlines():
            if line.startswith("event: "):
                event_type = line.removeprefix("event: ").strip()
            elif line.startswith("data: "):
                data_line = line.removeprefix("data: ").strip()
        if data_line is not None:
            events.append((event_type, json.loads(data_line)))
    return events


@pytest.mark.integration
def test_query_sse_event_sequence(client):
    response = client.post(
        "/api/v1/query/",
        json={
            "repo_id": "ad/adpilot-indexing-commit-intel.com",
            "question": "Explain the indexing pipeline end to end",
            "snapshot_id": "snap_demo",
            "options": {"stream": True},
        },
    )
    assert response.status_code == 200
    events = _parse_sse_events(response.text)
    event_types = [name for name, _ in events]
    assert event_types[0] == "start"
    assert "sources" in event_types
    assert event_types.count("token") >= 1
    assert event_types[-1] == "done"

    start_payload = events[0][1]
    assert start_payload["request_id"]
    assert start_payload["repo_id"] == "ad/adpilot-indexing-commit-intel.com"

    done_payload = events[-1][1]
    assert done_payload["answer_length"] > 0
    assert done_payload["answer"]
    assert len(done_payload["answer"]) == done_payload["answer_length"]
