from __future__ import annotations

import asyncio
import os
import socket
from collections.abc import Awaitable, Callable
from datetime import datetime, timezone
from typing import Any

from pydantic import ValidationError
from redis.asyncio import Redis
from redis.exceptions import ResponseError

from app.redis.retry import retry_while_redis_loading
from app.core.config import Settings, settings
from app.core.logger import logger
from app.models.stream_events import (
    STREAM_COMMIT_ANALYSIS_READY,
    STREAM_GRAPH_ARTIFACT_READY,
)
from app.models.validation import parse_commit_analysis_ready, parse_graph_artifact_ready
from app.services.stream_dispatch_service import (
    PermanentDispatchError,
    TransientDispatchError,
    dispatch_commit_analysis_ready,
    dispatch_graph_artifact_ready,
)

StreamHandler = Callable[[str], Awaitable[None]]


class StreamConsumer:
    """Redis Streams consumer for upstream embedding pipeline events."""

    def __init__(
        self,
        redis: Redis,
        app_settings: Settings | None = None,
        *,
        consumer_name: str | None = None,
    ) -> None:
        self.redis = redis
        self.settings = app_settings or settings
        self.consumer_name = consumer_name or self._resolve_consumer_name()
        self._handlers: dict[str, StreamHandler] = {
            STREAM_GRAPH_ARTIFACT_READY: self._handle_graph_artifact_ready,
            STREAM_COMMIT_ANALYSIS_READY: self._handle_commit_analysis_ready,
        }

    @staticmethod
    def _resolve_consumer_name() -> str:
        configured = settings.CONSUMER_NAME.strip()
        if configured:
            return configured
        hostname = socket.gethostname().strip()
        if hostname:
            return hostname
        return f"consumer-{os.getpid()}"

    def dlq_stream_name(self, source_stream: str) -> str:
        suffix = self.settings.STREAM_DLQ_SUFFIX.strip() or "dlq"
        return f"{source_stream}.{suffix}"

    async def ensure_group(self, stream: str) -> None:
        async def _create() -> None:
            try:
                await self.redis.xgroup_create(
                    stream,
                    self.settings.CONSUMER_GROUP,
                    id="$",
                    mkstream=True,
                )
                logger.info(
                    "Created consumer group group={group} stream={stream}",
                    group=self.settings.CONSUMER_GROUP,
                    stream=stream,
                )
            except ResponseError as exc:
                if "BUSYGROUP" not in str(exc):
                    raise

        await retry_while_redis_loading(
            f"ensure_group:{stream}",
            _create,
        )

    async def run(self, stop_event: asyncio.Event | None = None) -> None:
        stop_event = stop_event or asyncio.Event()
        tasks = [
            asyncio.create_task(
                self.consume_stream(base_stream, stop_event),
                name=f"consume-{base_stream}",
            )
            for base_stream in self._handlers
        ]
        try:
            await asyncio.gather(*tasks)
        finally:
            for task in tasks:
                task.cancel()
            await asyncio.gather(*tasks, return_exceptions=True)

    async def consume_stream(self, base_stream: str, stop_event: asyncio.Event) -> None:
        stream = self.settings.stream_name(base_stream)
        handler = self._handlers[base_stream]
        await self.ensure_group(stream)

        while not stop_event.is_set():
            try:
                batches = await self.redis.xreadgroup(
                    groupname=self.settings.CONSUMER_GROUP,
                    consumername=self.consumer_name,
                    streams={stream: ">"},
                    count=self.settings.STREAM_CONSUMER_BATCH_SIZE,
                    block=self.settings.STREAM_CONSUMER_BLOCK_MS,
                )
            except asyncio.CancelledError:
                raise
            except Exception as exc:
                logger.error(
                    "xreadgroup failed stream={stream} error={error}",
                    stream=stream,
                    error=str(exc),
                )
                await asyncio.sleep(1)
                continue

            if not batches:
                continue

            for _stream_name, messages in batches:
                for message_id, fields in messages:
                    await self.process_message(stream, message_id, fields, handler)

    async def process_message(
        self,
        stream: str,
        message_id: str,
        fields: dict[str, Any],
        handler: StreamHandler,
    ) -> None:
        delivery_count = await self._delivery_count(stream, message_id)
        payload = fields.get("payload")

        if not isinstance(payload, str) or not payload.strip():
            await self._move_to_dlq(
                stream,
                message_id,
                fields,
                reason="missing or empty payload field",
                delivery_count=delivery_count,
            )
            await self._ack(stream, message_id)
            return

        try:
            await handler(payload)
        except PermanentDispatchError as exc:
            await self._move_to_dlq(
                stream,
                message_id,
                fields,
                reason=str(exc),
                delivery_count=delivery_count,
            )
            await self._ack(stream, message_id)
            return
        except TransientDispatchError as exc:
            logger.error(
                "Transient stream handler failure stream={stream} message_id={message_id} error={error}",
                stream=stream,
                message_id=message_id,
                error=str(exc),
            )
            if delivery_count >= self.settings.STREAM_MAX_DELIVERY_ATTEMPTS:
                await self._move_to_dlq(
                    stream,
                    message_id,
                    fields,
                    reason=str(exc),
                    delivery_count=delivery_count,
                )
                await self._ack(stream, message_id)
            return
        except Exception as exc:
            logger.error(
                "Unexpected stream handler failure stream={stream} message_id={message_id} error={error}",
                stream=stream,
                message_id=message_id,
                error=str(exc),
            )
            if delivery_count >= self.settings.STREAM_MAX_DELIVERY_ATTEMPTS:
                await self._move_to_dlq(
                    stream,
                    message_id,
                    fields,
                    reason=str(exc),
                    delivery_count=delivery_count,
                )
                await self._ack(stream, message_id)
            return

        await self._ack(stream, message_id)

    async def _delivery_count(self, stream: str, message_id: str) -> int:
        try:
            pending = await self.redis.xpending_range(
                stream,
                self.settings.CONSUMER_GROUP,
                min=message_id,
                max=message_id,
                count=1,
            )
        except Exception:
            return 1

        if not pending:
            return 1

        entry = pending[0]
        if isinstance(entry, dict):
            return int(entry.get("times_delivered", 1))
        return int(getattr(entry, "times_delivered", 1))

    async def _ack(self, stream: str, message_id: str) -> None:
        try:
            await self.redis.xack(stream, self.settings.CONSUMER_GROUP, message_id)
        except Exception as exc:
            logger.error(
                "xack failed stream={stream} message_id={message_id} error={error}",
                stream=stream,
                message_id=message_id,
                error=str(exc),
            )

    async def _move_to_dlq(
        self,
        source_stream: str,
        message_id: str,
        fields: dict[str, Any],
        *,
        reason: str,
        delivery_count: int,
    ) -> None:
        dlq_stream = self.dlq_stream_name(source_stream)
        payload = {
            "source_stream": source_stream,
            "source_message_id": message_id,
            "reason": reason,
            "delivery_count": str(delivery_count),
            "failed_at": datetime.now(timezone.utc).isoformat(),
            "payload": fields.get("payload", ""),
        }
        try:
            await self.redis.xadd(dlq_stream, payload)
            logger.warning(
                "Moved message to DLQ dlq={dlq} source_stream={stream} message_id={message_id} reason={reason}",
                dlq=dlq_stream,
                stream=source_stream,
                message_id=message_id,
                reason=reason,
            )
        except Exception as exc:
            logger.error(
                "Failed to write DLQ entry dlq={dlq} message_id={message_id} error={error}",
                dlq=dlq_stream,
                message_id=message_id,
                error=str(exc),
            )

    async def _handle_graph_artifact_ready(self, payload: str) -> None:
        try:
            event = parse_graph_artifact_ready(payload)
        except (ValueError, ValidationError) as exc:
            raise PermanentDispatchError(f"invalid graph.artifact.ready payload: {exc}") from exc

        await dispatch_graph_artifact_ready(event)

    async def _handle_commit_analysis_ready(self, payload: str) -> None:
        try:
            event = parse_commit_analysis_ready(payload)
        except (ValueError, ValidationError) as exc:
            raise PermanentDispatchError(
                f"invalid commit.analysis.ready payload: {exc}"
            ) from exc

        await dispatch_commit_analysis_ready(event)
