import asyncio
from contextlib import suppress
from time import perf_counter
from typing import Optional

from redis.asyncio import Redis

from app.core.env import env_bool, env_float, env_int, env_str
from app.core.logger import logger
from app.services import embedding_job_service
from app.services.redis_consumer_service import run_files_changed_loop
from app.services import ingestion_queue_service
from app.services.stream_ingestion_service import (
    build_stream_key,
    process_file_changed_event,
)


class WorkerManager:
    """Bootstrap worker processes and stream consumers."""

    def __init__(self) -> None:
        self._task: Optional[asyncio.Task] = None
        self._redis: Optional[Redis] = None
        self._running: bool = False

    @property
    def running(self) -> bool:
        return self._running

    async def start(self) -> None:
        if self._running:
            return

        self._running = True

        redis_stream_enabled = env_bool("REDIS_STREAM_ENABLED", False)

        if redis_stream_enabled:
            self._task = asyncio.create_task(
                self._redis_stream_loop(),
                name="worker-redis-stream-loop",
            )
        else:
            self._task = asyncio.create_task(
                self._polling_loop(),
                name="worker-polling-loop",
            )

        logger.bind(
            concurrency=env_int("WORKER_CONCURRENCY", 2),
            redis_stream_enabled=redis_stream_enabled,
            queue_enabled=env_bool("QUEUE_ENABLED", False),
            queue_backend=env_str("QUEUE_BACKEND", "in_memory"),
            worker_task_name=self._task.get_name() if self._task else None,
        ).info("Worker manager started")

    async def stop(self) -> None:
        if not self._running:
            return

        self._running = False

        if self._task:
            self._task.cancel()

            with suppress(asyncio.CancelledError):
                await self._task

        if self._redis:
            await self._redis.aclose()
            self._redis = None

        self._task = None

        logger.info("Worker manager stopped")

    async def _polling_loop(self) -> None:
        poll_interval_seconds = env_float(
            "QUEUE_POLL_INTERVAL_SECONDS",
            1.0,
        )
        embedding_limit = env_int("EMBEDDING_JOB_DISPATCH_BATCH_SIZE", 100)
        ingestion_limit = env_int("QUEUE_BATCH_SIZE", 10)
        worker_concurrency = env_int("WORKER_CONCURRENCY", 2)

        logger.bind(
            queue_batch_size=ingestion_limit,
            worker_concurrency=worker_concurrency,
            poll_interval_seconds=poll_interval_seconds,
            embedding_batch_size=embedding_limit,
        ).info("Ingestion polling worker loop started")

        while self._running:
            loop_started = perf_counter()
            if env_bool("QUEUE_ENABLED", False):
                try:
                    ingestion_consumed = await ingestion_queue_service.consume_ingestion_jobs(
                        limit=ingestion_limit,
                        concurrency=worker_concurrency,
                    )
                    logger.bind(
                        ingestion_consumed=ingestion_consumed,
                        queue_batch_size=ingestion_limit,
                        worker_concurrency=worker_concurrency,
                    ).debug("Polling loop ingestion cycle completed")
                except Exception as exc:
                    logger.error(f"Ingestion queue worker loop failure: {exc}")

            if env_bool("EMBEDDING_JOB_WORKER_ENABLED", True):
                try:
                    await embedding_job_service.retry_failed_jobs(limit=embedding_limit)
                    await embedding_job_service.dispatch_embedding_jobs(limit=embedding_limit)
                except Exception as exc:
                    logger.error(f"Embedding job worker loop failure: {exc}")

            logger.bind(
                loop_duration_ms=round((perf_counter() - loop_started) * 1000, 2),
                queue_enabled=env_bool("QUEUE_ENABLED", False),
                embedding_enabled=env_bool("EMBEDDING_JOB_WORKER_ENABLED", True),
            ).debug("Polling loop iteration completed")

            await asyncio.sleep(poll_interval_seconds)

    async def _dispatch_embedding_jobs(self) -> None:
        if not env_bool("EMBEDDING_JOB_WORKER_ENABLED", True):
            return

        embedding_limit = env_int("EMBEDDING_JOB_DISPATCH_BATCH_SIZE", 100)
        try:
            await embedding_job_service.retry_failed_jobs(limit=embedding_limit)
            await embedding_job_service.dispatch_embedding_jobs(limit=embedding_limit)
        except Exception as exc:
            logger.error(f"Embedding job worker dispatch failure: {exc}")
        logger.info("Ingestion polling worker loop stopped")

    async def _redis_stream_loop(self) -> None:
        redis_url = env_str(
            "REDIS_URL",
            "redis://redis-server:6379",
        )

        self._redis = Redis.from_url(
            redis_url,
            decode_responses=True,
        )

        await self._redis.ping()

        await run_files_changed_loop(
            self._redis,
            running=lambda: self._running,
            on_idle=self._dispatch_embedding_jobs,
        )

        embedding_limit = env_int("EMBEDDING_JOB_DISPATCH_BATCH_SIZE", 100)
        ingestion_limit = env_int("QUEUE_BATCH_SIZE", 10)
        worker_concurrency = env_int("WORKER_CONCURRENCY", 2)

        while self._running:
            loop_started = perf_counter()
            try:
                stream_data = await self._redis.xread(
                    {stream_key: last_id},
                    block=block_ms,
                    count=count,
                )

                if stream_data:
                    stream_message_count = 0
                    for _, entries in stream_data:
                        stream_message_count += len(entries)
                        for message_id, fields in entries:
                            last_id = message_id

                            try:
                                event = parse_file_changed_event(fields)

                                await process_file_changed_event(event)

                            except Exception as exc:
                                logger.bind(
                                    stream_key=stream_key,
                                    message_id=message_id,
                                ).error(
                                    f"Failed processing stream event: {exc}"
                                )
                    logger.bind(
                        stream_key=stream_key,
                        stream_messages=stream_message_count,
                    ).debug("Redis stream batch processed")

                if env_bool("QUEUE_ENABLED", False):
                    try:
                        ingestion_consumed = await ingestion_queue_service.consume_ingestion_jobs(
                            limit=ingestion_limit,
                            concurrency=worker_concurrency,
                        )
                        logger.bind(
                            stream_key=stream_key,
                            ingestion_consumed=ingestion_consumed,
                            queue_batch_size=ingestion_limit,
                            worker_concurrency=worker_concurrency,
                        ).debug("Redis stream loop ingestion cycle completed")
                    except Exception as exc:
                        logger.error(f"Ingestion queue worker dispatch failure: {exc}")

                if env_bool("EMBEDDING_JOB_WORKER_ENABLED", True):
                    try:
                        await embedding_job_service.retry_failed_jobs(limit=embedding_limit)
                        await embedding_job_service.dispatch_embedding_jobs(limit=embedding_limit)
                    except Exception as exc:
                        logger.error(f"Embedding job worker dispatch failure: {exc}")

                logger.bind(
                    stream_key=stream_key,
                    loop_duration_ms=round((perf_counter() - loop_started) * 1000, 2),
                    queue_enabled=env_bool("QUEUE_ENABLED", False),
                    embedding_enabled=env_bool("EMBEDDING_JOB_WORKER_ENABLED", True),
                ).debug("Redis stream loop iteration completed")

            except asyncio.CancelledError:
                raise

            except Exception as exc:
                logger.bind(
                    stream_key=stream_key
                ).error(
                    f"Redis stream consumer error: {exc}"
                )

                await asyncio.sleep(poll_interval_seconds)

        logger.bind(stream_key=stream_key).info("Redis stream consumer stopped")


worker_manager = WorkerManager()
