import asyncio

from app.services import ingestion_queue_transports


def test_get_ingestion_queue_publisher_selects_backend():
    assert isinstance(
        ingestion_queue_transports.get_ingestion_queue_publisher("redis_stream"),
        ingestion_queue_transports.RedisStreamIngestionQueuePublisher,
    )
    assert isinstance(
        ingestion_queue_transports.get_ingestion_queue_publisher("kafka"),
        ingestion_queue_transports.KafkaIngestionQueuePublisher,
    )
    assert isinstance(
        ingestion_queue_transports.get_ingestion_queue_publisher("rabbitmq"),
        ingestion_queue_transports.RabbitMQIngestionQueuePublisher,
    )


def test_get_ingestion_queue_consumer_selects_backend():
    assert isinstance(
        ingestion_queue_transports.get_ingestion_queue_consumer("redis_stream"),
        ingestion_queue_transports.RedisStreamIngestionQueueConsumer,
    )
    assert isinstance(
        ingestion_queue_transports.get_ingestion_queue_consumer("kafka"),
        ingestion_queue_transports.KafkaIngestionQueueConsumer,
    )
    assert isinstance(
        ingestion_queue_transports.get_ingestion_queue_consumer("rabbitmq"),
        ingestion_queue_transports.RabbitMQIngestionQueueConsumer,
    )


def test_redis_consumer_acknowledges_processed_messages(monkeypatch):
    xacked = []

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

        async def xreadgroup(self, *args, **kwargs):
            return [("ingestion.jobs", [("1-0", {"payload": '{"run_id":"run_1"}'})])]

        async def xack(self, _stream, _group, message_id):
            xacked.append(message_id)

        async def aclose(self):
            return None

    async def _callback(_event):
        return None

    monkeypatch.setattr(ingestion_queue_transports.Redis, "from_url", lambda *_args, **_kwargs: _Redis())

    consumer = ingestion_queue_transports.RedisStreamIngestionQueueConsumer()
    consumed = asyncio.run(consumer.consume(_callback, limit=10))

    assert consumed == 1
    assert xacked == ["1-0"]


def test_redis_consumer_skips_ack_on_processing_failure(monkeypatch):
    xacked = []

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

        async def xreadgroup(self, *args, **kwargs):
            return [("ingestion.jobs", [("1-0", {"payload": '{"run_id":"run_1"}'})])]

        async def xack(self, _stream, _group, message_id):
            xacked.append(message_id)

        async def aclose(self):
            return None

    async def _callback(_event):
        raise RuntimeError("boom")

    monkeypatch.setattr(ingestion_queue_transports.Redis, "from_url", lambda *_args, **_kwargs: _Redis())

    consumer = ingestion_queue_transports.RedisStreamIngestionQueueConsumer()
    consumed = asyncio.run(consumer.consume(_callback, limit=10))

    assert consumed == 0
    assert xacked == []


def test_redis_consumer_processes_messages_concurrently(monkeypatch):
    xacked = []

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

        async def xreadgroup(self, *args, **kwargs):
            return [
                (
                    "ingestion.jobs",
                    [
                        ("1-0", {"payload": '{"run_id":"run_1"}'}),
                        ("2-0", {"payload": '{"run_id":"run_2"}'}),
                        ("3-0", {"payload": '{"run_id":"run_3"}'}),
                    ],
                )
            ]

        async def xack(self, _stream, _group, message_id):
            xacked.append(message_id)

        async def aclose(self):
            return None

    state = {"inflight": 0, "max_inflight": 0}

    async def _callback(_event):
        state["inflight"] += 1
        state["max_inflight"] = max(state["max_inflight"], state["inflight"])
        await asyncio.sleep(0)
        state["inflight"] -= 1

    monkeypatch.setattr(ingestion_queue_transports.Redis, "from_url", lambda *_args, **_kwargs: _Redis())

    consumer = ingestion_queue_transports.RedisStreamIngestionQueueConsumer()
    consumed = asyncio.run(consumer.consume(_callback, limit=10, concurrency=3))

    assert consumed == 3
    assert len(xacked) == 3
    assert state["max_inflight"] >= 2
