"""In-memory Event Store implementation.

Useful for testing and development. Events are lost on restart.
"""

from collections.abc import Iterator
from dataclasses import dataclass, field
import threading
from typing import Any, ClassVar

from mcp_hangar.domain.contracts.event_store import ConcurrencyError, IEventStore
from mcp_hangar.domain.events import DomainEvent
from mcp_hangar.domain.exceptions import CompactionError
from mcp_hangar.logging_config import get_logger

logger = get_logger(__name__)


@dataclass
class StoredEvent:
    """Event wrapper with metadata."""

    global_position: int
    stream_id: str
    stream_version: int
    event: DomainEvent


@dataclass
class Stream:
    """Stream state tracking."""

    stream_id: str
    version: int = -1
    events: list[StoredEvent] = field(default_factory=list)


class InMemoryEventStore(IEventStore):
    """In-memory event store for testing and development.

    Thread-safe but not persistent. All data is lost on restart.
    """

    #: One process, one lock, one list: a higher position was appended later.
    positions_are_commit_ordered: ClassVar[bool] = True

    def __init__(self):
        """Initialize empty event store."""
        self._streams: dict[str, Stream] = {}
        self._all_events: list[StoredEvent] = []
        self._snapshots: dict[str, dict[str, Any]] = {}
        # Re-entrant because `append_at_end` holds it while calling `append`,
        # which takes it again. Holding it across both keeps the version read and
        # the write one step.
        self._lock = threading.RLock()
        self._global_position = 0

        logger.info("in_memory_event_store_initialized")

    def append(
        self,
        stream_id: str,
        events: list[DomainEvent],
        expected_version: int,
    ) -> int:
        """Append events with optimistic concurrency."""
        if not events:
            return expected_version

        with self._lock:
            # Get or create stream
            stream = self._streams.get(stream_id)
            if stream is None:
                stream = Stream(stream_id=stream_id)
                self._streams[stream_id] = stream

            # Check version
            if stream.version != expected_version:
                raise ConcurrencyError(stream_id, expected_version, stream.version)

            # Append events
            for event in events:
                self._global_position += 1
                stream.version += 1

                stored = StoredEvent(
                    global_position=self._global_position,
                    stream_id=stream_id,
                    stream_version=stream.version,
                    event=event,
                )
                stream.events.append(stored)
                self._all_events.append(stored)

            logger.debug(
                "events_appended",
                stream_id=stream_id,
                events_count=len(events),
                new_version=stream.version,
            )

            return stream.version

    def append_at_end(self, stream_id: str, events: list[DomainEvent]) -> int:
        """Append after whatever the stream holds, under the lock every writer takes.

        The version is read and the batch appended within one hold of
        `self._lock`, which `append` takes too. No other writer can move the
        stream in between, so a batch that claimed no version cannot conflict.
        The lock is re-entrant so that this can go
        through `append` rather than beside it.
        """
        if not events:
            return self.get_stream_version(stream_id)
        with self._lock:
            return self.append(stream_id, events, self.get_stream_version(stream_id))

    def read_stream(
        self,
        stream_id: str,
        from_version: int = 0,
    ) -> list[DomainEvent]:
        """Read events from stream."""
        with self._lock:
            stream = self._streams.get(stream_id)
            if stream is None:
                return []

            return [stored.event for stored in stream.events if stored.stream_version >= from_version]

    def read_all(
        self,
        from_position: int = 0,
        limit: int = 1000,
    ) -> Iterator[tuple[int, str, DomainEvent]]:
        """Read all events globally."""
        with self._lock:
            events = [e for e in self._all_events if e.global_position > from_position][:limit]

        for stored in events:
            yield stored.global_position, stored.stream_id, stored.event

    def get_stream_version(self, stream_id: str) -> int:
        """Get current stream version."""
        with self._lock:
            stream = self._streams.get(stream_id)
            return stream.version if stream else -1

    def clear(self) -> None:
        """Clear all events (for testing)."""
        with self._lock:
            self._streams.clear()
            self._all_events.clear()
            self._global_position = 0

        logger.info("event_store_cleared")

    def get_event_count(self) -> int:
        """Get total event count."""
        with self._lock:
            return len(self._all_events)

    def get_stream_count(self) -> int:
        """Get total stream count."""
        with self._lock:
            return len(self._streams)

    def list_streams(self, prefix: str = "") -> list[str]:
        """List all stream IDs, optionally filtered by prefix."""
        with self._lock:
            if prefix:
                return [sid for sid in self._streams.keys() if sid.startswith(prefix)]
            return list(self._streams.keys())

    def save_snapshot(
        self,
        stream_id: str,
        version: int,
        state: dict[str, Any],
    ) -> None:
        """Save aggregate snapshot in memory."""
        with self._lock:
            self._snapshots[stream_id] = {
                "version": version,
                "state": state,
            }
            logger.debug("snapshot_saved", stream_id=stream_id, version=version)

    def load_snapshot(
        self,
        stream_id: str,
    ) -> dict[str, Any] | None:
        """Load latest snapshot for a stream."""
        with self._lock:
            return self._snapshots.get(stream_id)

    def compact_stream(self, stream_id: str) -> int:
        """Delete events that precede the latest snapshot for a stream.

        Args:
            stream_id: Identifier of the stream to compact.

        Returns:
            Number of events deleted.

        Raises:
            CompactionError: When no snapshot exists for the stream.
        """
        with self._lock:
            snapshot = self._snapshots.get(stream_id)
            if snapshot is None:
                raise CompactionError(stream_id, "no snapshot exists; create a snapshot before compacting")

            snapshot_version: int = snapshot["version"]
            stream = self._streams.get(stream_id)

            if stream is None:
                return 0

            before_count = len(stream.events)
            compacted = [e for e in stream.events if e.stream_version > snapshot_version]
            deleted = before_count - len(compacted)
            stream.events = compacted

            # Also remove from global list
            self._all_events = [
                e for e in self._all_events if not (e.stream_id == stream_id and e.stream_version <= snapshot_version)
            ]

        from ...metrics import record_events_compacted

        record_events_compacted(stream_id, deleted)

        logger.info(
            "stream_compacted",
            stream_id=stream_id,
            snapshot_version=snapshot_version,
            events_deleted=deleted,
        )

        return deleted
