import pytest

from app.services import embedding_publishers


@pytest.mark.parametrize(
    "backend,expected",
    [
        ("redis_stream", embedding_publishers.RedisStreamEmbeddingPublisher),
        ("redis", embedding_publishers.RedisStreamEmbeddingPublisher),
        ("kafka", embedding_publishers.KafkaEmbeddingPublisher),
        ("rabbitmq", embedding_publishers.RabbitMQEmbeddingPublisher),
    ],
)
def test_get_embedding_publisher_selects_provider(backend, expected):
    publisher = embedding_publishers.get_embedding_publisher(backend)
    assert isinstance(publisher, expected)


def test_get_embedding_publisher_rejects_unknown_backend():
    with pytest.raises(RuntimeError):
        embedding_publishers.get_embedding_publisher("unknown")


@pytest.mark.asyncio
async def test_redis_publisher_wraps_connection_errors(monkeypatch):
    calls = []

    class _BindLogger:
        def error(self, msg):
            calls.append(msg)

    class _Logger:
        def bind(self, **_kwargs):
            return _BindLogger()

    class _Redis:
        async def xadd(self, *_args, **_kwargs):
            raise RuntimeError("socket timeout")

        async def aclose(self):
            return None

    monkeypatch.setattr(embedding_publishers, "logger", _Logger())
    monkeypatch.setattr(embedding_publishers.Redis, "from_url", lambda *_args, **_kwargs: _Redis())

    publisher = embedding_publishers.RedisStreamEmbeddingPublisher()

    with pytest.raises(RuntimeError, match="Queue connection failure \\(redis\\)"):
        await publisher.publish({"job_id": "emb_1"})

    assert any("Queue connection failure (redis)" in msg for msg in calls)
