"""Provider registry — pluggable chat / completion backends.

Providers are pure Python objects implementing :class:`ChatProvider`. The
``EchoProvider`` ships as a built-in deterministic backend so the rest of
the AI layer is testable without any external service.
"""

from __future__ import annotations

from collections.abc import Iterable
from dataclasses import dataclass, field
from typing import Any, Protocol


class ProviderError(RuntimeError):
    """Raised on registration errors or provider invocation failures."""


@dataclass(frozen=True)
class Message:
    role: str          # "system" | "user" | "assistant" | "tool"
    content: str
    name: str = ""

    def __post_init__(self) -> None:
        if self.role not in {"system", "user", "assistant", "tool"}:
            raise ProviderError(f"unknown message role {self.role!r}")


@dataclass(frozen=True)
class ChatResponse:
    content: str
    provider: str
    model: str
    finish_reason: str = "stop"
    usage: dict[str, int] = field(default_factory=dict)


class ChatProvider(Protocol):
    name: str

    def chat(self, messages: Iterable[Message], **opts: Any) -> ChatResponse: ...


class EchoProvider:
    """Deterministic offline provider.

    Concatenates the last user message with a short prefix. Used by tests
    and as a default when no real provider is configured.
    """

    name = "echo"

    def __init__(self, model: str = "echo-1") -> None:
        self.model = model

    def chat(self, messages: Iterable[Message], **opts: Any) -> ChatResponse:
        msgs = list(messages)
        last_user = next(
            (m.content for m in reversed(msgs) if m.role == "user"),
            "",
        )
        content = f"[echo] {last_user}"
        return ChatResponse(
            content=content,
            provider=self.name,
            model=self.model,
            usage={
                "prompt_tokens": sum(len(m.content.split()) for m in msgs),
                "completion_tokens": len(content.split()),
            },
        )


class ScriptedProvider:
    """Reads canned replies in order — handy for asserting agent flow in tests."""

    name = "scripted"

    def __init__(self, replies: list[str], model: str = "scripted-1") -> None:
        self._replies = list(replies)
        self._idx = 0
        self.model = model

    def chat(self, _messages: Iterable[Message], **_opts: Any) -> ChatResponse:
        if self._idx >= len(self._replies):
            raise ProviderError("scripted provider exhausted")
        content = self._replies[self._idx]
        self._idx += 1
        return ChatResponse(content=content, provider=self.name, model=self.model)


_PROVIDERS: dict[str, ChatProvider] = {}


def register_provider(provider: ChatProvider) -> ChatProvider:
    name = getattr(provider, "name", "")
    if not name:
        raise ProviderError("provider must expose a non-empty .name")
    _PROVIDERS[name] = provider
    return provider


def get_provider(name: str) -> ChatProvider:
    try:
        return _PROVIDERS[name]
    except KeyError as exc:
        raise ProviderError(f"unknown provider {name!r}") from exc


def list_providers() -> list[str]:
    return sorted(_PROVIDERS)


def reset_for_tests() -> None:
    _PROVIDERS.clear()
    register_default_providers()


def register_default_providers() -> None:
    if "echo" not in _PROVIDERS:
        register_provider(EchoProvider())
