"""Dependency resolver for fixture providers using topological sort.

Given a set of FixtureProviders with depends_on specifications, resolves the
correct execution order. Detects cycles and provides diagnostic information.
"""

from __future__ import annotations

from collections import defaultdict, deque
from typing import TYPE_CHECKING

if TYPE_CHECKING:
    from simorgh.apps.provisioning.providers import FixtureProvider


class CyclicDependencyError(Exception):
    """Raised when fixture providers have circular dependencies."""

    def __init__(self, cycle: list[str]) -> None:
        self.cycle = cycle
        cycle_str = " -> ".join(cycle)
        super().__init__(f"Cyclic fixture dependency detected: {cycle_str}")


class UnresolvedDependencyError(Exception):
    """Raised when a fixture depends on an unknown/unregistered provider."""

    def __init__(self, provider: str, dependency: str) -> None:
        self.provider = provider
        self.dependency = dependency
        super().__init__(
            f"Provider {provider!r} depends on {dependency!r} "
            f"which is not registered."
        )


class DependencyResolver:
    """Resolves execution order for fixture providers.

    Uses Kahn's algorithm for topological sorting. Providers with no
    dependencies execute first. Providers at the same dependency level
    execute in registration order.

    Usage:
        resolver = DependencyResolver()
        resolver.register(provider_a)  # depends_on=()
        resolver.register(provider_b)  # depends_on=("provider_a",)
        order = resolver.resolve()     # [provider_a, provider_b]
        groups = resolver.resolve_grouped()  # [[provider_a], [provider_b]]
    """

    def __init__(self) -> None:
        self._providers: dict[str, FixtureProvider] = {}
        self._graph: dict[str, set[str]] = defaultdict(set)
        self._reverse_graph: dict[str, set[str]] = defaultdict(set)

    def register(self, provider: FixtureProvider) -> None:
        """Register a fixture provider with the resolver."""
        name = provider.name
        if name in self._providers:
            raise ValueError(f"Provider {name!r} is already registered.")
        self._providers[name] = provider
        self._graph[name] = set()
        self._reverse_graph[name] = set()

    def register_bulk(self, providers: list[FixtureProvider]) -> None:
        for p in providers:
            self.register(p)

    def _build_edges(self) -> None:
        """Build dependency edges after all providers are registered."""
        for name, provider in self._providers.items():
            for dep in provider.depends_on:
                if dep not in self._providers:
                    raise UnresolvedDependencyError(name, dep)
                self._graph[dep].add(name)
                self._reverse_graph[name].add(dep)

    def resolve(self) -> list[FixtureProvider]:
        """Return providers in topological execution order.

        Raises:
            CyclicDependencyError: If a cycle is detected.
            UnresolvedDependencyError: If a dependency is not registered.
        """
        self._build_edges()

        in_degree = {name: len(deps) for name, deps in self._reverse_graph.items()}
        queue = deque(name for name, deg in in_degree.items() if deg == 0)
        order: list[FixtureProvider] = []

        while queue:
            name = queue.popleft()
            order.append(self._providers[name])
            for dependent in sorted(self._graph[name]):
                in_degree[dependent] -= 1
                if in_degree[dependent] == 0:
                    queue.append(dependent)

        if len(order) != len(self._providers):
            remaining = [n for n, d in in_degree.items() if d > 0]
            cycle = self._detect_cycle(remaining[0])
            raise CyclicDependencyError(cycle)

        return order

    def resolve_grouped(self) -> list[list[FixtureProvider]]:
        """Return providers grouped by dependency level for parallel execution."""
        self._build_edges()

        in_degree = {name: len(deps) for name, deps in self._reverse_graph.items()}
        current = [name for name, deg in in_degree.items() if deg == 0]
        groups: list[list[FixtureProvider]] = []

        while current:
            group = [self._providers[name] for name in sorted(current)]
            groups.append(group)
            next_level: list[str] = []
            for name in current:
                for dependent in sorted(self._graph[name]):
                    in_degree[dependent] -= 1
                    if in_degree[dependent] == 0:
                        next_level.append(dependent)
            current = next_level

        if sum(len(g) for g in groups) != len(self._providers):
            remaining = [n for n, d in in_degree.items() if d > 0]
            cycle = self._detect_cycle(remaining[0])
            raise CyclicDependencyError(cycle)

        return groups

    def resolve_names(self) -> list[str]:
        """Return ordered list of provider names."""
        return [p.name for p in self.resolve()]

    def _detect_cycle(self, start: str) -> list[str]:
        """Trace a cycle from the given starting node using DFS."""
        visited: set[str] = set()
        path: list[str] = []
        cycle: list[str] = []

        def dfs(node: str) -> bool:
            if node in visited:
                if node in path:
                    cycle.extend(path[path.index(node):])
                    cycle.append(node)
                    return True
                return False
            visited.add(node)
            path.append(node)
            for neighbor in self._reverse_graph.get(node, set()):
                if dfs(neighbor):
                    return True
            path.pop()
            return False

        dfs(start)
        return cycle if cycle else [start]

    def reset(self) -> None:
        self._providers.clear()
        self._graph.clear()
        self._reverse_graph.clear()

    @property
    def provider_count(self) -> int:
        return len(self._providers)

    def list_providers(self) -> dict[str, list[str]]:
        """Return a summary of registered providers and their dependencies."""
        return {
            name: list(provider.depends_on)
            for name, provider in self._providers.items()
        }

    def reverse_order(self) -> list[FixtureProvider]:
        """Return providers in reverse-dependency order (for resets)."""
        return list(reversed(self.resolve()))


def resolve_fixtures(
    providers: list[FixtureProvider],
) -> list[FixtureProvider]:
    """Convenience function to resolve a list of providers."""
    resolver = DependencyResolver()
    resolver.register_bulk(providers)
    return resolver.resolve()
