"""Discover AFL games, load every available odds row, and follow each game.

Install: python -m pip install httpx
Run:     ODDS_API_KEY=your_key python afl_full_feed.py

The program writes JSON Lines to stdout. Set AFL_PAST_HOURS and
AFL_FUTURE_HOURS to change the rolling discovery window.
"""

from __future__ import annotations

import asyncio
import json
import os
import random
import sys
import time
from typing import Any, AsyncIterator

import httpx


BASE_URL = os.getenv("ODDS_API_BASE_URL", "https://api.odds-api.net/v1").rstrip("/") + "/"
API_KEY = os.getenv("ODDS_API_KEY", "").strip()
PAST_HOURS = int(os.getenv("AFL_PAST_HOURS", "8"))
FUTURE_HOURS = int(os.getenv("AFL_FUTURE_HOURS", "48"))
DISCOVERY_SECONDS = int(os.getenv("AFL_DISCOVERY_SECONDS", "300"))
ODDS_SHAPE = {
    "price_fields": "all",
    "include_unavailable": "true",
    "include_source": "true",
}


def emit(kind: str, **payload: Any) -> None:
    print(json.dumps({"kind": kind, **payload}, separators=(",", ":")), flush=True)


async def get_json(client: httpx.AsyncClient, path: str, params: dict | None = None) -> dict:
    response = await client.get(path, params=params)
    response.raise_for_status()
    return response.json()


async def afl_events(client: httpx.AsyncClient) -> list[dict]:
    now = int(time.time())
    params = {
        "sport": "australian rules",
        "league": "AFL",
        "start_from": now - PAST_HOURS * 3600,
        "start_to": now + FUTURE_HOURS * 3600,
        "limit": 1000,
    }
    events: list[dict] = []
    while True:
        page = await get_json(client, "events", params)
        events.extend(page.get("items") or [])
        cursor = page.get("next_cursor")
        if not cursor:
            return events
        params["cursor"] = cursor


async def complete_snapshot(client: httpx.AsyncClient, event_id: str) -> tuple[dict, dict[str, dict]]:
    params = {"limit": 10000, **ODDS_SHAPE}
    rows: dict[str, dict] = {}
    first_page: dict | None = None
    while True:
        page = await get_json(client, f"events/{event_id}/odds/snapshot", params)
        if first_page is None:
            first_page = page
        for odd in page.get("items") or []:
            if odd.get("id"):
                rows[str(odd["id"])] = odd
        cursor = page.get("next_cursor")
        if not cursor:
            break
        params["cursor"] = cursor

    # Start the stream at the first page's resume token. Changes that arrive
    # while later snapshot pages load can then be replayed by the stream.
    return first_page or {}, rows


async def sse_messages(response: httpx.Response) -> AsyncIterator[tuple[str, dict]]:
    event = "message"
    data: list[str] = []
    async for line in response.aiter_lines():
        if line == "":
            if data:
                yield event, json.loads("\n".join(data))
            event, data = "message", []
        elif line.startswith("event:"):
            event = line[6:].strip()
        elif line.startswith("data:"):
            data.append(line[5:].lstrip())
    if data:
        yield event, json.loads("\n".join(data))


def apply_delta(rows: dict[str, dict], payload: dict) -> None:
    for change in payload.get("changes") or []:
        odd = change.get("odd") or {}
        row_id = str(odd.get("id") or "")
        if not row_id:
            continue
        if change.get("op") in {"remove", "delete"}:
            rows.pop(row_id, None)
        elif change.get("op") == "upsert":
            # Keep is_available=false rows so the product can show suspensions.
            rows[row_id] = odd


async def follow_event(client: httpx.AsyncClient, event_id: str) -> None:
    resume = ""
    rows: dict[str, dict] = {}
    needs_snapshot = True
    retry_seconds = 1.0

    while True:
        try:
            if needs_snapshot:
                snapshot, rows = await complete_snapshot(client, event_id)
                resume = str(snapshot.get("resume") or "$")
                emit(
                    "odds_snapshot",
                    event_id=event_id,
                    as_of_ts_ms=snapshot.get("as_of_ts_ms"),
                    ttl_seconds=snapshot.get("ttl_seconds"),
                    resume=resume,
                    items=list(rows.values()),
                )
                needs_snapshot = False

            params = {"since": resume, "catchup": "true", "max_batch": 2000, **ODDS_SHAPE}
            async with client.stream(
                "GET", f"events/{event_id}/odds/stream", params=params,
                headers={"Accept": "text/event-stream"},
            ) as response:
                response.raise_for_status()
                retry_seconds = 1.0
                async for event, payload in sse_messages(response):
                    if event == "delta":
                        apply_delta(rows, payload)
                        resume = str(payload.get("resume") or resume)
                        emit("odds_delta", event_id=event_id, resume=resume, changes=payload.get("changes") or [])
                    elif event == "resync":
                        emit("odds_resync", event_id=event_id, reason=payload.get("reason"))
                        needs_snapshot = True
                        break
                    # A heartbeat contains no price change.

        except asyncio.CancelledError:
            raise
        except (httpx.HTTPError, ValueError, json.JSONDecodeError) as exc:
            emit("stream_error", event_id=event_id, error=str(exc))

        if needs_snapshot:
            continue
        await asyncio.sleep(retry_seconds + random.uniform(0, 0.5))
        retry_seconds = min(retry_seconds * 2, 30.0)


async def game_record(client: httpx.AsyncClient, event_id: str) -> None:
    try:
        detail, result = await asyncio.gather(
            get_json(client, f"events/{event_id}"),
            get_json(client, f"events/{event_id}/results"),
        )
        emit("event", event_id=event_id, data=detail)
        emit("result", event_id=event_id, data=result)
    except httpx.HTTPError as exc:
        emit("record_error", event_id=event_id, error=str(exc))


async def bounded_game_record(client: httpx.AsyncClient, semaphore: asyncio.Semaphore, event_id: str) -> None:
    async with semaphore:
        await game_record(client, event_id)


async def main() -> None:
    if not API_KEY:
        raise SystemExit("Set ODDS_API_KEY before running this example.")
    timeout = httpx.Timeout(connect=10.0, read=None, write=10.0, pool=10.0)
    tasks: dict[str, asyncio.Task] = {}
    record_limit = asyncio.Semaphore(6)
    async with httpx.AsyncClient(
        base_url=BASE_URL,
        headers={"X-API-Key": API_KEY, "Accept": "application/json"},
        timeout=timeout,
    ) as client:
        try:
            while True:
                try:
                    events = await afl_events(client)
                    ids = {str(item["event_id"]) for item in events if item.get("event_id")}
                    emit("event_list", count=len(events), items=events)

                    for event_id in ids:
                        if event_id not in tasks or tasks[event_id].done():
                            tasks[event_id] = asyncio.create_task(follow_event(client, event_id))
                    for event_id in set(tasks) - ids:
                        tasks.pop(event_id).cancel()

                    await asyncio.gather(
                        *(bounded_game_record(client, record_limit, event_id) for event_id in ids)
                    )
                except httpx.HTTPError as exc:
                    emit("discovery_error", error=str(exc))
                await asyncio.sleep(DISCOVERY_SECONDS)
        finally:
            for task in tasks.values():
                task.cancel()
            await asyncio.gather(*tasks.values(), return_exceptions=True)


if __name__ == "__main__":
    try:
        asyncio.run(main())
    except KeyboardInterrupt:
        print("Stopped.", file=sys.stderr)
