import asyncio
import json

import pytest

from app.models.enums import EmbedStatus
from app.models.stream_events import (
    CommitAnalysisReadyEvent,
    GraphArtifactReadyEvent,
)
from app.redis.stream_consumer import StreamConsumer
from app.services import stream_dispatch_service


GRAPH_PAYLOAD = json.dumps(
    {
        "event_id": "evt_snap_abc123_graph_ready",
        "event_version": "v1",
        "repo_id": "AD/example-repo",
        "snapshot_id": "snap_abc123",
        "commit_sha": "a1b2c3d4e5f6",
        "artifact_id": "art_snap_abc123_merged",
        "artifact_uri": "mongo://snapshot_graphs/art_snap_abc123_merged",
    }
)

COMMIT_PAYLOAD = json.dumps(
    {
        "event_id": "evt_analysis_01",
        "event_version": "v1",
        "repo_id": "AD/example-repo",
        "commit_sha": "b2c3d4e5f6a7",
        "analysis_id": "analysis_b2c3d4",
        "summary": "Introduced Reload",
    }
)


class _FakeRedis:
    def __init__(self):
        self.acked: list[tuple[str, str, str]] = []
        self.dlq_entries: list[tuple[str, dict]] = []
        self.pending_deliveries: dict[str, int] = {}
        self.xreadgroup_result: list = []
        self.xreadgroup_calls = 0

    async def xgroup_create(self, *_args, **_kwargs):
        return None

    async def xreadgroup(self, **_kwargs):
        self.xreadgroup_calls += 1
        if self.xreadgroup_calls == 1:
            return self.xreadgroup_result
        await asyncio.sleep(0.01)
        return []

    async def xack(self, stream, group, message_id):
        self.acked.append((stream, group, message_id))

    async def xpending_range(self, stream, group, min, max, count):
        return [{"message_id": min, "times_delivered": self.pending_deliveries.get(min, 1)}]

    async def xadd(self, stream, payload):
        self.dlq_entries.append((stream, payload))


@pytest.fixture
def graph_event():
    return GraphArtifactReadyEvent.model_validate(json.loads(GRAPH_PAYLOAD))


def test_dispatch_graph_artifact_enqueues_celery(monkeypatch, graph_event):
    task_ids: list[str] = []

    class _Task:
        id = "celery-task-1"

    def _delay(*args, **kwargs):
        task_ids.append(args)
        return _Task()

    monkeypatch.setattr(
        stream_dispatch_service.process_snapshot_embedding_task,
        "delay",
        _delay,
    )

    async def _upsert_run(**kwargs):
        from app.models.embedding_run import EmbeddingRun, build_run_dedupe_key, build_run_id

        dedupe = build_run_dedupe_key(kwargs["repo_id"], kwargs["snapshot_id"], kwargs["event_id"])
        return EmbeddingRun(
            run_id=build_run_id(dedupe),
            dedupe_key=dedupe,
            repo_id=kwargs["repo_id"],
            snapshot_id=kwargs["snapshot_id"],
            event_id=kwargs["event_id"],
            run_type=kwargs.get("run_type", "snapshot_batch"),
        )

    async def _get_run_by_dedupe_key(_key):
        return None

    async def _update_run_state(*_args, **_kwargs):
        return True

    monkeypatch.setattr(stream_dispatch_service.embedding_run_repo, "upsert_run", _upsert_run)
    monkeypatch.setattr(
        stream_dispatch_service.embedding_run_repo,
        "get_run_by_dedupe_key",
        _get_run_by_dedupe_key,
    )
    monkeypatch.setattr(
        stream_dispatch_service.embedding_run_repo,
        "update_run_state",
        _update_run_state,
    )

    result = asyncio.run(stream_dispatch_service.dispatch_graph_artifact_ready(graph_event))

    assert result.skipped is False
    assert result.celery_task_id == "celery-task-1"
    assert task_ids[0][0] == result.run_id


def test_dispatch_skips_duplicate_event(monkeypatch, graph_event):
    from app.models.embedding_run import EmbeddingRun

    existing = EmbeddingRun(
        run_id="erun_existing",
        dedupe_key="dedupe",
        repo_id=graph_event.repo_id,
        snapshot_id=graph_event.snapshot_id,
        event_id=graph_event.event_id,
        status=EmbedStatus.PENDING,
        celery_task_id="already-enqueued",
    )

    async def _get_run_by_dedupe_key(_key):
        return existing

    monkeypatch.setattr(
        stream_dispatch_service.embedding_run_repo,
        "get_run_by_dedupe_key",
        _get_run_by_dedupe_key,
    )
    monkeypatch.setattr(
        stream_dispatch_service.process_snapshot_embedding_task,
        "delay",
        lambda *_args, **_kwargs: pytest.fail("should not enqueue duplicate"),
    )

    result = asyncio.run(stream_dispatch_service.dispatch_graph_artifact_ready(graph_event))
    assert result.skipped is True
    assert result.celery_task_id == "already-enqueued"


def test_consumer_acks_valid_graph_event(monkeypatch):
    fake_redis = _FakeRedis()
    fake_redis.xreadgroup_result = [
        ("graph.artifact.ready", [("1-0", {"payload": GRAPH_PAYLOAD})])
    ]

    async def _dispatch(_event):
        return stream_dispatch_service.DispatchResult(
            run_id="erun_test",
            celery_task_id="task-1",
        )

    monkeypatch.setattr(
        "app.redis.stream_consumer.dispatch_graph_artifact_ready",
        _dispatch,
    )

    consumer = StreamConsumer(fake_redis)
    stop_event = asyncio.Event()

    async def _stop_soon():
        await asyncio.sleep(0.05)
        stop_event.set()

    async def _run():
        await asyncio.gather(
            consumer.consume_stream("graph.artifact.ready", stop_event),
            _stop_soon(),
        )

    asyncio.run(_run())

    assert fake_redis.acked == [("graph.artifact.ready", "embedding-engine-service", "1-0")]


def test_consumer_moves_invalid_payload_to_dlq():
    fake_redis = _FakeRedis()
    fake_redis.xreadgroup_result = [
        ("graph.artifact.ready", [("2-0", {"payload": "{bad-json"})])
    ]

    consumer = StreamConsumer(fake_redis)
    stop_event = asyncio.Event()

    async def _stop_soon():
        await asyncio.sleep(0.05)
        stop_event.set()

    async def _run():
        await asyncio.gather(
            consumer.consume_stream("graph.artifact.ready", stop_event),
            _stop_soon(),
        )

    asyncio.run(_run())

    assert fake_redis.acked == [("graph.artifact.ready", "embedding-engine-service", "2-0")]
    assert len(fake_redis.dlq_entries) == 1
    assert fake_redis.dlq_entries[0][0] == "graph.artifact.ready.dlq"


def test_consumer_does_not_ack_transient_failure(monkeypatch):
    fake_redis = _FakeRedis()
    fake_redis.xreadgroup_result = [
        ("commit.analysis.ready", [("3-0", {"payload": COMMIT_PAYLOAD})])
    ]

    async def _dispatch(_event):
        raise stream_dispatch_service.TransientDispatchError("broker unavailable")

    monkeypatch.setattr(
        "app.redis.stream_consumer.dispatch_commit_analysis_ready",
        _dispatch,
    )

    consumer = StreamConsumer(fake_redis)
    stop_event = asyncio.Event()

    async def _stop_soon():
        await asyncio.sleep(0.05)
        stop_event.set()

    async def _run():
        await asyncio.gather(
            consumer.consume_stream("commit.analysis.ready", stop_event),
            _stop_soon(),
        )

    asyncio.run(_run())

    assert fake_redis.acked == []


def test_consumer_dlq_after_max_delivery_attempts(monkeypatch):
    fake_redis = _FakeRedis()
    fake_redis.pending_deliveries["4-0"] = 5
    fake_redis.xreadgroup_result = [
        ("graph.artifact.ready", [("4-0", {"payload": GRAPH_PAYLOAD})])
    ]

    async def _dispatch(_event):
        raise stream_dispatch_service.TransientDispatchError("still failing")

    monkeypatch.setattr(
        "app.redis.stream_consumer.dispatch_graph_artifact_ready",
        _dispatch,
    )

    consumer = StreamConsumer(fake_redis)
    stop_event = asyncio.Event()

    async def _stop_soon():
        await asyncio.sleep(0.05)
        stop_event.set()

    async def _run():
        await asyncio.gather(
            consumer.consume_stream("graph.artifact.ready", stop_event),
            _stop_soon(),
        )

    asyncio.run(_run())

    assert fake_redis.acked == [("graph.artifact.ready", "embedding-engine-service", "4-0")]
    assert fake_redis.dlq_entries


def test_invalid_graph_payload_raises_permanent_error():
    with pytest.raises(stream_dispatch_service.PermanentDispatchError):
        asyncio.run(
            StreamConsumer(_FakeRedis())._handle_graph_artifact_ready('{"event_id":"only-id"}')
        )
