"""Relay subscriber for the persisted Hermes shared-metrics slice."""

from __future__ import annotations

import logging
import platform
import threading
from typing import Any

from agent.relay_runtime import RUNTIME_INSTANCE_KEY
from hermes_cli.config import detect_install_method
from hermes_constants import get_hermes_home

from . import shared_metrics_engagement as engagement
from . import shared_metrics_signals as signals
from .shared_metrics import SharedMetricsStore
from .shared_metrics_fields import milestones_for
from .shared_metrics_contract import (
    CLIENT_ACTIVE_METRIC,
    COMMIT_TICKET_KEY,
    FEATURE_DISABLED_METRIC,
    FEATURE_USED_MARK,
    INSTALL_SNAPSHOT_METRIC,
    MODEL_ROUTE_METRIC,
    TOOL_CALL_METRIC,
    TOOL_LATENCY_METRIC,
    TOOL_USAGE_METRIC,
    client_active_counter,
    client_resource,
    decision_counter,
    install_snapshot_counter,
    model_call_dimensions,
    model_token_counters,
    skill_counter,
    task_counter,
    task_duration_counter,
    tool_approval_counter,
    tool_call_dimensions,
    tool_latency_dimensions,
    tool_usage_dimensions,
)

logger = logging.getLogger(__name__)

# Contract projections; each yields (metric_name, dimensions) or None. One tool end event feeds
# the category, per-tool and latency counters (one task end event the task and duration counters),
# so every match is recorded.
_COUNTERS = (
    client_active_counter,
    install_snapshot_counter,
    lambda event: _named(MODEL_ROUTE_METRIC, model_call_dimensions(event)),
    lambda event: _named(TOOL_CALL_METRIC, tool_call_dimensions(event)),
    lambda event: _named(TOOL_USAGE_METRIC, tool_usage_dimensions(event)),
    lambda event: _named(TOOL_LATENCY_METRIC, tool_latency_dimensions(event)),
    task_counter,
    task_duration_counter,
    tool_approval_counter,
    skill_counter,
    decision_counter,
    signals.feature_used_counter,
)


def _named(metric_name: str, dimensions: dict | None) -> tuple[str, dict] | None:
    return None if dimensions is None else (metric_name, dimensions)


class SharedMetricsSubscriber:
    """Persist validated Hermes counters from Relay lifecycle events."""

    def __init__(
        self,
        store: SharedMetricsStore,
        hermes_version: str,
        *,
        runtime_id: str | None = None,
    ) -> None:
        self.store = store
        self._client_resource = client_resource(
            hermes_version,
            os_name=platform.system(),
            architecture=platform.machine(),
            install_method=detect_install_method(),
        )
        self._runtime_id = runtime_id
        self._active = True
        self._lock = threading.RLock()
        self._milestones_done: set[str] = set(store.recorded_milestones())
        self._features_done: set[str] = set(store.recorded_features())
        self._saved_tickets: dict[str, int] = {}
        # Events arrive on the Relay thread, which carries no profile binding.
        self._hermes_home = get_hermes_home()

    def deactivate(self) -> None:
        """Stop accepting events before telemetry is disabled or torn down."""
        with self._lock:
            self._active = False

    @staticmethod
    def _classify(event: Any) -> list[tuple[str, dict, int]]:
        """Return every ``(metric_name, dimensions, amount)`` the event satisfies."""
        counted = [(*m, 1) for m in (project(event) for project in _COUNTERS) if m is not None]
        return counted + model_token_counters(event)

    def _record_milestones(self, metric_name: str, dimensions: dict) -> None:
        """Latch every install milestone this counter reaches (each once per install, ever)."""
        reached = [m for m in milestones_for(metric_name, dimensions) if m not in self._milestones_done]
        if not reached:
            return
        from .shared_metrics_snapshot import install_age_bucket

        age = install_age_bucket(self._hermes_home)
        for milestone in reached:
            self.store.record_milestone(milestone, age, self._client_resource)
            self._milestones_done.add(milestone)

    def _record_engagement(self, event: Any, classified: list[tuple[str, dict, int]]) -> None:
        """Fold a turn start/end or an attended turn's model into the local daily engagement rollup."""
        surface, route = engagement.interaction_surface(classified), engagement.turn_route(event)
        if surface is None and route is None:
            return
        with self._lock:
            if not self._active:
                return
            try:
                engagement.record(self.store, self._client_resource, surface=surface, route=route)
            except Exception:
                logger.warning("Unable to update the Hermes engagement rollup", exc_info=True)

    def _persist(self, metric_name: str, dimensions: dict, amount: int) -> None:
        store, resource = self.store, self._client_resource
        special = {
            FEATURE_USED_MARK: lambda: None,  # a first-use fact, not a counter row
            FEATURE_DISABLED_METRIC: lambda: store.record_counter_once_per_day(metric_name, dimensions, resource),
            CLIENT_ACTIVE_METRIC: lambda: store.record_client_active(resource),
            INSTALL_SNAPSHOT_METRIC: lambda: store.record_install_snapshot(dimensions, resource),
        }.get(metric_name)
        if special is not None:
            special()
        else:
            store.record_counter(metric_name, dimensions, resource, amount)

    def _record_features(self, metric_name: str, dimensions: dict) -> None:
        """Latch each feature's first real use (once per install, with the owning profile's age)."""
        reached = [f for f in signals.features_for(metric_name, dimensions) if f not in self._features_done]
        if not reached:
            return
        age = signals.days_since_install_bucket(self._hermes_home)
        for feature in reached:
            try:
                self.store.record_feature_adoption(feature, age, self._client_resource)
            except Exception:  # the counter row is saved; a later use retries the latch
                logger.debug("Feature adoption not latched: %s", feature, exc_info=True)
                continue
            self._features_done.add(feature)

    def take_saved(self, ticket: str) -> int:
        """How many events carrying ``ticket`` settled without a store error (and forget the ticket)."""
        with self._lock:
            return self._saved_tickets.pop(ticket, 0)

    def __call__(self, event: Any) -> None:
        metadata = getattr(event, "metadata", None)
        if self._runtime_id is not None:
            if (
                not isinstance(metadata, dict)
                or metadata.get(RUNTIME_INSTANCE_KEY) != self._runtime_id
            ):
                return
        ticket = metadata.get(COMMIT_TICKET_KEY) if isinstance(metadata, dict) else None
        saved = True  # a row the contract rejects is settled too: no retry can change that
        classified = self._classify(event)
        self._record_engagement(event, classified)
        for metric_name, dimensions, amount in classified:
            with self._lock:
                if not self._active:
                    return
                try:
                    self._persist(metric_name, dimensions, amount)
                    self._record_milestones(metric_name, dimensions)
                    self._record_features(metric_name, dimensions)
                except Exception:
                    saved = False
                    logger.warning(
                        "Unable to persist the Hermes shared metric: %s", metric_name, exc_info=True
                    )
        if ticket and saved:
            with self._lock:
                if self._active:
                    self._saved_tickets[ticket] = self._saved_tickets.get(ticket, 0) + 1
