import asyncio

from app.workers import bootstrap


class _RedisStub:
    def __init__(self, manager):
        self._manager = manager

    async def ping(self):
        return True

    async def xread(self, *_args, **_kwargs):
        # Stop after one iteration to make test deterministic.
        self._manager._running = False
        return []

    async def aclose(self):
        return None


def test_polling_loop_forwards_worker_concurrency(monkeypatch):
    manager = bootstrap.WorkerManager()
    manager._running = True

    called = {"limit": None, "concurrency": None}

    def _env_bool(name, default=False):
        if name == "QUEUE_ENABLED":
            return True
        if name == "EMBEDDING_JOB_WORKER_ENABLED":
            return False
        return default

    def _env_int(name, default=0):
        if name == "QUEUE_BATCH_SIZE":
            return 7
        if name == "WORKER_CONCURRENCY":
            return 4
        return default

    async def _consume_ingestion_jobs(limit, concurrency):
        called["limit"] = limit
        called["concurrency"] = concurrency
        manager._running = False
        return 1

    async def _sleep(_seconds):
        return None

    monkeypatch.setattr(bootstrap, "env_bool", _env_bool)
    monkeypatch.setattr(bootstrap, "env_int", _env_int)
    monkeypatch.setattr(bootstrap, "env_float", lambda *_args, **_kwargs: 0.0)
    monkeypatch.setattr(bootstrap.ingestion_queue_service, "consume_ingestion_jobs", _consume_ingestion_jobs)
    monkeypatch.setattr(bootstrap.asyncio, "sleep", _sleep)

    asyncio.run(manager._polling_loop())

    assert called["limit"] == 7
    assert called["concurrency"] == 4


def test_redis_stream_loop_polls_ingestion_queue_when_stream_is_idle(monkeypatch):
    manager = bootstrap.WorkerManager()
    manager._running = True

    called = {"ingestion": 0}

    def _env_bool(name, default=False):
        if name == "QUEUE_ENABLED":
            return True
        if name == "EMBEDDING_JOB_WORKER_ENABLED":
            return False
        return default

    def _env_int(name, default=0):
        if name == "QUEUE_BATCH_SIZE":
            return 5
        if name == "WORKER_CONCURRENCY":
            return 3
        return default

    async def _consume_ingestion_jobs(limit, concurrency):
        called["ingestion"] += 1
        assert limit == 5
        assert concurrency == 3
        return 0

    monkeypatch.setattr(bootstrap, "env_bool", _env_bool)
    monkeypatch.setattr(bootstrap, "env_int", _env_int)
    monkeypatch.setattr(bootstrap, "env_float", lambda *_args, **_kwargs: 0.0)
    monkeypatch.setattr(bootstrap, "env_str", lambda *_args, **_kwargs: "redis://redis-server:6379")
    monkeypatch.setattr(bootstrap, "build_stream_key", lambda: "files.changed")
    monkeypatch.setattr(bootstrap.Redis, "from_url", lambda *_args, **_kwargs: _RedisStub(manager))
    monkeypatch.setattr(bootstrap.ingestion_queue_service, "consume_ingestion_jobs", _consume_ingestion_jobs)

    asyncio.run(manager._redis_stream_loop())

    assert called["ingestion"] == 1


def test_worker_manager_start_stop_lifecycle(monkeypatch):
    manager = bootstrap.WorkerManager()

    async def _loop(self):
        while self._running:
            await asyncio.sleep(0)

    def _env_bool(name, default=False):
        if name == "REDIS_STREAM_ENABLED":
            return False
        return default

    monkeypatch.setattr(bootstrap.WorkerManager, "_polling_loop", _loop)
    monkeypatch.setattr(bootstrap, "env_bool", _env_bool)
    monkeypatch.setattr(bootstrap, "env_int", lambda *_args, **_kwargs: 2)
    monkeypatch.setattr(bootstrap, "env_str", lambda *_args, **_kwargs: "redis_stream")

    async def _run():
        await manager.start()
        assert manager.running is True
        assert manager._task is not None
        await manager.stop()
        assert manager.running is False
        assert manager._task is None

    asyncio.run(_run())
