"""Direct NeMo Relay integration for Hermes shared client metrics."""

from __future__ import annotations

import atexit
import contextlib
import contextvars
import logging
import threading
import uuid
from collections import OrderedDict, deque
from dataclasses import dataclass, field
from time import monotonic_ns
from typing import Any, Callable

from agent import relay_runtime
from agent.portal_tags import get_conversation_context
from hermes_cli.version_info import get_version_info

from .shared_metrics import SharedMetricsStore
from . import shared_metrics_contract as contract
from . import shared_metrics_efficiency as eff
from . import shared_metrics_engagement as engagement_
from . import shared_metrics_fields as fields_
from . import shared_metrics_model as model_
from .shared_metrics_contract import MODEL_CALL_SCOPE, SUBSCRIBER_NAME, TASK_SCOPE
from .shared_metrics_subscriber import SharedMetricsSubscriber

logger = logging.getLogger(__name__)

_RUNTIME_FAILED = object()
_RUNTIMES: dict[str, _Runtime | object] = {}
_RUNTIME_LOCK = threading.RLock()

_ABORTED = {"failed": True, "turn_exit_reason": "system_aborted"}
# The store latch allows one snapshot per 24h; re-reading it hourly keeps task starts cheap.
_SNAPSHOT_RECHECK_NS = 3_600 * 1_000_000_000


def _text(event: dict[str, Any], key: str) -> str:
    return str(event.get(key) or "")


def _session_pair(event: dict[str, Any], key: str) -> tuple[str, str] | None:
    """(session_id, event[key]) when both are non-empty."""
    session_id, value = _text(event, "session_id"), _text(event, key)
    return (session_id, value) if session_id and value else None


def _retry_ordinal(event: dict[str, Any]) -> int:
    """Hermes's provider-local retry ordinal; 0 when absent or malformed."""
    value = event.get("retry_count")
    return value if isinstance(value, int) and not isinstance(value, bool) and value > 0 else 0


def _forget(index: dict[Any, _MetricsSession], key: Any, owner: _MetricsSession) -> None:
    """Drop ``key`` from ``index`` only while it still points at ``owner``."""
    if index.get(key) is owner:
        index.pop(key, None)


def _task_parent_handle(session: _MetricsSession, task_id: str) -> Any:
    """The active turn's handle when it owns this exact task, else the session handle."""
    active_turn = relay_runtime.active_turn(session.session_id)
    if (
        active_turn is not None
        and active_turn.lease.session_id == session.session_id
        and active_turn.task_id == task_id
        and active_turn.handle is not None
    ):
        return active_turn.handle
    return session.relay_session.handle


def _elapsed_ms(started_ns: int) -> int:
    return max(0, (monotonic_ns() - started_ns) // 1_000_000)


def _scope_handle(session: _MetricsSession, task: _TaskRun | None) -> Any:
    return task.handle if task is not None else session.relay_session.handle


def _sole(items: Any) -> Any:
    """The single distinct element of ``items`` (identity-deduplicated), else None."""
    unique = {id(item): item for item in items}
    return next(iter(unique.values())) if len(unique) == 1 else None


def _identities_compatible(candidate: tuple[str, str, str], observed: tuple[str, str, str]) -> bool:
    """Match partial hook context without crossing known call boundaries."""
    if not observed[2] or candidate[2] != observed[2]:
        return False
    return all(
        not left or not right or left == right
        for left, right in zip(candidate[:2], observed[:2], strict=True)
    )


def _compatible_tool_call_keys(
    session: _MetricsSession, task_id: str, identity: tuple[str, str, str]
) -> list[tuple[str, str, str, str]]:
    return [
        key
        for key in session.tool_calls
        if key[0] == task_id and _identities_compatible(key[1:], identity)
    ]


@dataclass
class _ModelCall:
    handle: Any
    task_id: str
    fields: dict[str, str]
    error_class: str = "none"
    ttft_bucket: str = "unknown"
    started_ns: int = field(default_factory=monotonic_ns)


@dataclass
class _ToolCall:
    handle: Any
    category: str
    tool_name: str
    started_ns: int
    approval_outcome: str = "not_required"


@dataclass
class _TaskRun:
    task_id: str
    handle: Any
    context: contextvars.Context
    started_ns: int
    start_fields: dict[str, str]
    model_call_ids: set[str] = field(default_factory=set)
    tool_call_ids: set[tuple[str, str, str]] = field(default_factory=set)
    turn_ids: set[str] = field(default_factory=set)
    retired_turn_ids: frozenset[str] = field(default_factory=frozenset)
    completed_tool_call_ids: set[tuple[str, str, str]] = field(default_factory=set)
    unidentified_tool_calls: int = 0
    retry_count: int = 0
    model_route: dict[str, str] | None = None
    # The route the turn was sent on (its first request): provider failover does not change it.
    selected_route: dict[str, str] | None = None
    cost: eff.TurnCost = field(default_factory=eff.TurnCost)


@dataclass
class _MetricsSession:
    session_id: str
    relay_session: relay_runtime.RelaySession
    lock: threading.RLock = field(default_factory=threading.RLock, repr=False)
    closing: bool = False
    model_calls: dict[tuple[str, str], _ModelCall] = field(default_factory=dict)
    tasks: dict[str, _TaskRun] = field(default_factory=dict)
    tool_calls: dict[tuple[str, str, str, str], _ToolCall] = field(default_factory=dict)
    retired_turn_ids: deque[str] = field(default_factory=lambda: deque(maxlen=256))
    # Session-level aggregation, emitted once as hermes.session.count when the session closes.
    start_fields: dict[str, str] | None = None
    turns: int = 0
    failed_turns: int = 0
    last_outcome: str = "unknown"
    first_turn_ns: int = 0
    last_turn_ns: int = 0
    model_state: model_.ModelSessionState = field(default_factory=model_.ModelSessionState)
    # Compression rotated this id away: close once its in-flight turn ends.
    retiring: bool = False
    efficiency: eff.SessionEfficiency = field(default_factory=eff.SessionEfficiency)
    route_run: engagement_.RouteRun = field(default_factory=engagement_.RouteRun)
    api_calls: int = 0
    tool_call_total: int = 0
    replies: int = 0


@dataclass
class _PeakLineage:
    """Open segments of one conversation (compression rotates the session id), their merged peak and
    session tally, the conversation's current-model turn run and its spent user turns."""

    open: set[str] = field(default_factory=set)
    peak: model_.ModelSessionState = field(default_factory=model_.ModelSessionState)
    tools: eff.ToolUsage = field(default_factory=eff.ToolUsage)
    tally: engagement_.SessionTally = field(default_factory=engagement_.SessionTally)
    run: engagement_.RouteRun = field(default_factory=engagement_.RouteRun)
    spent: deque | None = None


# Backstop: conversations whose rotated-to segment a surface never closed are flushed past this many.
_MAX_LINEAGES = 512


def _absorb_peak(into: model_.ModelSessionState, segment: model_.ModelSessionState) -> None:
    """Merge one closed segment's context peak (fullest fill wins, a limit hit anywhere sticks)."""
    if segment.peak_route is None:
        return
    into.limit_hit = into.limit_hit or segment.limit_hit
    fuller = segment.peak_window is not None and (
        into.peak_window is None or segment.peak_tokens / segment.peak_window > into.peak_tokens / into.peak_window
    )
    if into.peak_route is None or fuller:
        into.peak_route, into.peak_tokens, into.peak_window = segment.peak_route, segment.peak_tokens, segment.peak_window


class _Runtime:
    """Own shared-metrics state layered on the Hermes core Relay host."""

    def __init__(self, host: relay_runtime.RelayRuntime | None = None) -> None:
        resolved_host = host or relay_runtime.get_runtime()
        if resolved_host is None:
            raise RuntimeError("Hermes core Relay runtime is unavailable")
        self.host: relay_runtime.RelayRuntime = resolved_host
        self.relay = self.host.relay
        self._active = True
        self._sessions: dict[str, _MetricsSession] = {}
        self._task_sessions: dict[tuple[str, str], _MetricsSession] = {}
        self._turn_sessions: dict[tuple[str, str], _MetricsSession] = {}
        self._sessions_lock = threading.RLock()
        # Leaf lock (nothing is acquired under it): taken while a session.lock is held. Lineages are
        # keyed by the conversation's first segment id; _lineage_of maps every open segment id to it.
        self._lineages: dict[str, _PeakLineage] = {}
        self._lineage_of: dict[str, str] = {}
        self._lineage_lock = threading.Lock()
        # Leaf lock: conversations whose next cold cache read Hermes already announced.
        self._cold_expected: OrderedDict[str, None] = OrderedDict()
        self._cold_lock = threading.Lock()
        self._task_creation_lock = threading.RLock()
        self._task_sessions_lock = threading.RLock()
        # Guards the opt-in send pass: at most one in flight per process.
        self._send_lock = threading.RLock()
        self._send_thread: threading.Thread | None = None
        self._snapshot_checked_ns: int | None = None
        self._subscriber_name = f"{SUBSCRIBER_NAME}.{self.host.runtime_id}"
        self.subscriber = SharedMetricsSubscriber(
            SharedMetricsStore(), get_version_info().base_version, runtime_id=self.host.runtime_id
        )
        self.relay.subscribers.register(self._subscriber_name, self.subscriber)
        self.host.retain_managed_execution(self._subscriber_name)
        self._registered = True
        atexit.register(self.shutdown)

    def ensure_session(self, event: dict[str, Any]) -> _MetricsSession | None:
        session_id = _text(event, "session_id")
        if not session_id:
            return None
        with self._sessions_lock:
            if not self._active:
                return None
            relay_session = self.host.ensure_session(event)
            if relay_session is None:
                return None
            session = self._sessions.get(session_id)
            if session is None:
                session = _MetricsSession(session_id=session_id, relay_session=relay_session)
                self._sessions[session_id] = session
                self._share_spent(session)
        with session.lock:
            return None if session.closing else session

    def record_client_active(self, event: dict[str, Any]) -> None:
        """Emit one payload-free activation attempt under the session scope."""
        session = self.ensure_session(event)
        if session is not None:
            self._emit_client_active(session)

    def _emit_client_active(self, session: _MetricsSession) -> None:
        with session.lock:
            if not session.closing:
                self._mark(session, None, contract.CLIENT_ACTIVE_MARK, {})
                self._safe(self._emit_install_snapshot_if_due, session)

    def _emit_install_snapshot_if_due(self, session: _MetricsSession) -> None:
        now = monotonic_ns()
        if self._snapshot_checked_ns is not None and now - self._snapshot_checked_ns < _SNAPSHOT_RECHECK_NS:
            return
        self._snapshot_checked_ns = now
        if not self.subscriber.store.install_snapshot_due():
            return
        from .shared_metrics_snapshot import collect_install_snapshot

        fields = collect_install_snapshot(_raw_config())
        self._mark(session, None, contract.INSTALL_SNAPSHOT_MARK, fields)

    def _mark(
        self, session: _MetricsSession, task: _TaskRun | None, name: str, data: dict[str, str]
    ) -> None:
        """Emit one Relay mark under the task scope when given, else the session scope."""
        self._run_scoped(
            session, task, self.relay.scope.event, name,
            handle=_scope_handle(session, task), data=data, metadata=self._event_metadata(),
        )

    def start_task(self, event: dict[str, Any]) -> _TaskRun | None:
        """Open one Relay function scope for a Hermes task run."""
        task_key = _session_pair(event, "task_id")
        if task_key is None:
            return None
        _, task_id = task_key
        with self._task_creation_lock:
            owner = self._task_session(event)
            if owner is not None:
                with owner.lock:
                    if owner.closing:
                        return None
                    task = owner.tasks.get(task_id)
                    if task is not None and not self._admits(owner, task, event):
                        return None
                    return task

            session = self.ensure_session(event)
            if session is None:
                return None
            with session.lock:
                turn_id = _text(event, "turn_id")
                if (
                    session.closing
                    or (turn_id and turn_id in session.retired_turn_ids)
                    or session.relay_session.context is None
                ):
                    return None
                self._emit_client_active(session)
                task_context = session.relay_session.context.copy()
                start_fields = contract.task_start_fields(event)
                handle = task_context.run(
                    self._with_scope_stack, self.relay.scope.push,
                    TASK_SCOPE, self.relay.ScopeType.Function,
                    handle=_task_parent_handle(session, task_id), input=start_fields,
                    metadata=self._event_metadata(),
                )
                task = _TaskRun(
                    task_id=task_id,
                    handle=handle,
                    context=task_context,
                    started_ns=monotonic_ns(),
                    start_fields=start_fields,
                    retired_turn_ids=frozenset(session.retired_turn_ids),
                )
                session.tasks[task_id] = task
                with self._task_sessions_lock:
                    self._task_sessions[task_key] = session
                self._remember_turn(session, task, event)
                return task

    def start_user_turn(self, event: dict[str, Any]) -> None:
        """pre_llm_call: a user message started this task (Hermes-owned forks never fire it)."""
        task = self.start_task(event)
        if task is not None:
            task.cost.user_turn = True

    def _run_in_task(
        self, task: _TaskRun, callback: Callable[..., Any], *args: Any, **kwargs: Any
    ) -> Any:
        return task.context.copy().run(self._with_scope_stack, callback, *args, **kwargs)

    def _with_scope_stack(self, callback: Callable[..., Any], *args: Any, **kwargs: Any) -> Any:
        self.relay.get_scope_stack()
        return callback(*args, **kwargs)

    def start_model_call(self, event: dict[str, Any]) -> None:
        task_id = _text(event, "task_id")
        session, task = self._task_pair(event, start=True, allow_task_id_fallback=True)
        if task_id and task is None:
            return
        session = session or self.ensure_session(event)
        if session is None:
            return
        request_id = _text(event, "api_request_id")
        if not request_id:
            return
        model_call_key = (task_id, request_id)
        fields = contract.model_call_fields(event)
        with session.lock:
            if session.closing:
                return
            if task is not None and not self._admits(session, task, event, current=True):
                return
            existing = session.model_calls.get(model_call_key)
            if existing is not None:
                existing.fields = fields
                if task is not None:
                    # Every repeated start for one logical request is another physical
                    # attempt. Provider fallback resets Hermes's provider-local retry
                    # ordinal, so ordinal deltas are not a reliable task-level counter.
                    task.retry_count += 1
                return
            if task is not None:
                task.model_route = fields
                task.selected_route = task.selected_route or fields
                task.model_call_ids.add(request_id)
                if _retry_ordinal(event) > 0:
                    # A real Hermes retry can advance api_request_id while carrying the
                    # retry ordinal. Count that physical attempt.
                    task.retry_count += 1
            handle = self._run_scoped(
                session, task, self.relay.llm.call, MODEL_CALL_SCOPE, self.relay.LLMRequest({}, {}),
                handle=_scope_handle(session, task), metadata=self._event_metadata(),
                model_name=contract.MODEL_CALL_PROFILE_MODEL,
            )
            session.model_calls[model_call_key] = _ModelCall(handle, task_id, fields)

    def update_model_call(self, event: dict[str, Any], *, finish: bool) -> None:
        """Refresh the located model call's fields from ``event``; ``finish`` closes it.

        ``api_request_error`` retains the latest attempt error without closing the logical
        call; ``post_api_request`` closes it.
        """
        session = self._any_session(event)
        if session is None:
            return
        with session.lock:
            if session.closing:
                return
            model_call_key = self._existing_model_call_key(session, event)
            model_call = session.model_calls.get(model_call_key) if model_call_key else None
            if model_call is None:
                return
            model_call.fields = contract.model_call_fields(event)
            if finish:
                task = session.tasks.get(model_call.task_id)
                session.replies += task is not None and task.cost.user_turn
                session.model_state.observe_call(model_call.fields, event.get("usage"), event.get("context_length"))
                self._observe_turn_call(session, session.tasks.get(model_call.task_id), model_call, event)
                model_call.ttft_bucket = fields_.ttft_bucket(event)
                self._finish_model_call(session, model_call_key, "success")
                tokens = fields_.model_token_fields(
                    event.get("usage"), call_role="primary", **model_call.fields
                )
                if tokens is not None:
                    self._guarded(
                        "Hermes shared-metrics token mark failed", self._mark,
                        session, session.tasks.get(model_call.task_id), contract.MODEL_TOKENS_MARK, tokens,
                    )
            else:
                model_call.error_class = contract.model_error_class(event)
                session.model_state.observe_error(model_call.fields, model_call.error_class)

    def start_tool_call(self, event: dict[str, Any]) -> None:
        """Open one privacy-safe Relay tool lifecycle under its task."""
        task_id = _text(event, "task_id")
        session, task = self._task_pair(event, start=True, allow_task_id_fallback=True)
        if session is None or task is None or not _text(event, "tool_call_id"):
            return
        identity = self._tool_call_identity(event)
        with session.lock:
            if not self._admits(session, task, event):
                return
            key = (task_id, *identity)
            if identity in task.completed_tool_call_ids or key in session.tool_calls:
                return
            task.tool_call_ids.add(identity)
            session.tool_calls[key] = self._open_tool_call(task, event)

    def record_approval(self, event: dict[str, Any]) -> None:
        """Record one bounded approval result without approval text or commands."""
        session, task = self._approval_task(event)
        if session is None or task is None:
            return
        outcome = contract.tool_approval_outcome(event)
        attribution = "unattributed"
        with session.lock:
            if session.closing or not self._event_matches_task_turn(task, event):
                return
            if _text(event, "tool_call_id"):
                identity = self._tool_call_identity(event)
                tool_call = session.tool_calls.get((task.task_id, *identity))
                if tool_call is None:
                    key = _sole(_compatible_tool_call_keys(session, task.task_id, identity))
                    tool_call = session.tool_calls[key] if key is not None else None
                if tool_call is not None:
                    tool_call.approval_outcome = outcome
                    attribution = "tool_call"
            self._mark(
                session, task, contract.TOOL_APPROVAL_MARK,
                {"attribution": attribution, "outcome": outcome},
            )

    def record_tool_call(self, event: dict[str, Any]) -> None:
        """Close and count one unique privacy-safe tool lifecycle."""
        task_id = _text(event, "task_id")
        session, task = self._task_pair(event, allow_task_id_fallback=True)
        if session is None or task is None:
            return
        with session.lock:
            if not self._admits(session, task, event):
                return
            tool_call = None
            if _text(event, "tool_call_id"):
                observed_identity = self._tool_call_identity(event)
                if observed_identity in task.completed_tool_call_ids:
                    return
                identity = observed_identity
                tool_call = session.tool_calls.pop((task_id, *identity), None)
                if tool_call is None:
                    if any(
                        _identities_compatible(completed, observed_identity)
                        for completed in task.completed_tool_call_ids
                    ):
                        return
                    matching_keys = _compatible_tool_call_keys(session, task_id, observed_identity)
                    if len(matching_keys) > 1:
                        # Partial context cannot safely choose between concurrent calls
                        # that reused the provider-local ID.
                        return
                    if matching_keys:
                        identity = matching_keys[0][1:]
                        tool_call = session.tool_calls.pop(matching_keys[0])
                task.completed_tool_call_ids.update({identity, observed_identity})
                task.tool_call_ids.add(identity)
            else:
                task.unidentified_tool_calls += 1
            if tool_call is None:
                tool_call = self._open_tool_call(task, event)
            if task.cost.user_turn:
                session.efficiency.tools.used.add(eff.toolset_metric_name(event.get("toolset")))
            self._finish_tool_call(task, tool_call, event)

    def record_skill_lifecycle(self, event: dict[str, Any]) -> None:
        """Emit one allowlisted skill fact without its local identity."""
        if _text(event, "action").strip().lower() == "loaded":
            mark, fields = contract.SKILL_LOAD_MARK, contract.skill_load_fields(event)
        else:
            mark, fields = contract.SKILL_LIFECYCLE_MARK, contract.skill_lifecycle_fields(event)
        if fields is None:
            return

        session_id, task_id = _text(event, "session_id"), _text(event, "task_id")
        session, task = self._task_pair(event, allow_task_id_fallback=not session_id)
        if session is None:
            if not (session_id and task_id):
                # No owning task: a bare process-level mark.
                self._with_scope_stack(
                    self.relay.scope.event, mark, data=fields, metadata=self._event_metadata()
                )
            return
        if task is None:
            return
        with session.lock:
            if (
                not session.closing
                and session.tasks.get(task.task_id) is task
                and self._event_matches_task_turn(task, event)
            ):
                self._mark(session, task, mark, fields)

    def finish_task(self, event: dict[str, Any]) -> None:
        """Close one task scope exactly once with bounded terminal fields."""
        session = self._any_session(event)
        if session is None:
            return
        with session.lock:
            finished = not session.closing and self._finish_task(
                session, _text(event, "task_id"), event
            )
            retired = finished and session.retiring and not session.tasks
        if retired:
            self.close_session({"session_id": session.session_id})
        elif finished:
            self._flush_and_export("Hermes shared-metrics task flush failed")

    def close_session(self, event: dict[str, Any]) -> None:
        session = self._session(event)
        if session is None:
            self._close_unseen_segment(_text(event, "session_id"))
            return
        if not self._abort_session(
            session, {**event, **_ABORTED, "completed": False, "interrupted": False}
        ):
            return
        self._emit_session_summary(session)
        try:
            self.relay.subscribers.flush()
        except Exception as exc:
            logger.warning(
                "Hermes shared-metrics session %s closed with errors: subscriber flush failed: %s",
                session.session_id,
                exc,
            )
        else:
            self._export()
        with self._sessions_lock:
            _forget(self._sessions, session.session_id, session)

    def shutdown(self) -> None:
        with self._sessions_lock:
            self._active = False
            session_ids = list(self._sessions)
        for session_id in session_ids:
            self._safe(self.close_session, {"session_id": session_id})
        with self._lineage_lock:
            unseen = list(self._lineage_of)
        for session_id in unseen:
            self._safe(self._close_unseen_segment, session_id)
        if not self._registered:
            return
        self._flush_and_export("Hermes shared-metrics shutdown flush failed")
        self._deregister()
        self._release()

    def _deregister(self) -> None:
        self._safe(self.relay.subscribers.deregister, self._subscriber_name)
        self.host.release_managed_execution(self._subscriber_name)
        self._registered = False

    def deactivate(self) -> None:
        """Stop collection without exporting locally aggregated metrics."""
        with self._sessions_lock:
            self._active = False
        self.subscriber.deactivate()
        if self._registered:
            self._deregister()
        with self._sessions_lock:
            sessions = list(self._sessions.values())
        for session in sessions:
            self._abort_session(session, {"session_id": session.session_id, **_ABORTED})
        with self._sessions_lock:
            self._sessions.clear()
        with self._task_sessions_lock:
            self._task_sessions.clear()
            self._turn_sessions.clear()
        with self._lineage_lock:
            self._lineages.clear()
            self._lineage_of.clear()
        self._release()

    def _release(self) -> None:
        """Let an in-flight send finish briefly, then drop the atexit hook.

        A short-lived CLI process would otherwise exit and kill the daemon send thread
        mid-request — the common case for this feature's one cadence.
        """
        self._join_send_thread()
        with contextlib.suppress(Exception):
            atexit.unregister(self.shutdown)

    def _join_send_thread(self, timeout: float = 2.0) -> None:
        """Bounded on purpose: pending packages stay in SQLite and go out next run, so
        blocking on a slow network is the wrong trade; the daemon thread dies with the process."""
        with self._send_lock:
            thread = self._send_thread
        if thread is not None and thread.is_alive():
            try:
                thread.join(timeout)
            except Exception:
                logger.debug("Shared-metrics send thread join failed", exc_info=True)

    def _session(self, event: dict[str, Any]) -> _MetricsSession | None:
        with self._sessions_lock:
            return self._sessions.get(_text(event, "session_id"))

    def _any_session(self, event: dict[str, Any]) -> _MetricsSession | None:
        """Owner session by task/turn correlation, else by session_id."""
        return self._task_session(event, allow_task_id_fallback=True) or self._session(event)

    def _task_pair(
        self, event: dict[str, Any], *, start: bool = False, **lookup: Any
    ) -> tuple[_MetricsSession | None, _TaskRun | None]:
        """Resolve (session, task) for a task-scoped hook; ``start`` opens a missing task."""
        session = self._task_session(event, **lookup)
        task = session.tasks.get(_text(event, "task_id")) if session is not None else None
        if task is None and start:
            task = self.start_task(event)
            session = self._task_session(event) if task is not None else None
        return session, task

    def _run_scoped(
        self, session: _MetricsSession, task: _TaskRun | None, callback: Callable[..., Any],
        *args: Any, **kwargs: Any,
    ) -> Any:
        """Run under the task context when the call belongs to a task, else the session."""
        if task is not None:
            return self._run_in_task(task, callback, *args, **kwargs)
        return self.host.run_in_session(session.relay_session, callback, *args, **kwargs)

    def _flush_and_export(self, failure_message: str) -> None:
        """Flush the Relay subscriber, then export; a failed flush skips the export."""
        try:
            self.relay.subscribers.flush()
        except Exception:
            logger.warning(failure_message, exc_info=True)
        else:
            self._export()

    def _abort_session(self, session: _MetricsSession, base_event: dict[str, Any]) -> bool:
        """Mark the session closing and system-abort its open tasks; False if already closing."""
        with session.lock:
            if session.closing:
                return False
            session.closing = True
            for task_id in list(session.tasks):
                self._finish_task(session, task_id, {**base_event, "task_id": task_id})
            self._end_pending_model_calls(session, base_event)
        return True

    def _task_session(
        self, event: dict[str, Any], *, allow_task_id_fallback: bool = False
    ) -> _MetricsSession | None:
        """Owner session by (session, turn), then (session, task), then unique task_id."""
        session_id, task_id = _text(event, "session_id"), _text(event, "task_id")
        if not task_id:
            return None
        with self._task_sessions_lock:
            owner = self._turn_sessions.get(_session_pair(event, "turn_id"))
            if owner is None and session_id:
                owner = self._task_sessions.get((session_id, task_id))
            if owner is not None or not allow_task_id_fallback:
                return owner
            return _sole(
                session for (_, tid), session in self._task_sessions.items() if tid == task_id
            )

    def _remember_turn(
        self, session: _MetricsSession, task: _TaskRun, event: dict[str, Any]
    ) -> None:
        turn_id = _text(event, "turn_id")
        if turn_id:
            task.turn_ids.add(turn_id)
            with self._task_sessions_lock:
                self._turn_sessions[(session.session_id, turn_id)] = session

    @staticmethod
    def _tool_call_identity(event: dict[str, Any]) -> tuple[str, str, str]:
        """Identify one provider-local tool call without exporting its IDs."""
        return _text(event, "api_request_id"), _text(event, "turn_id"), _text(event, "tool_call_id")

    @staticmethod
    def _event_matches_task_turn(task: _TaskRun, event: dict[str, Any]) -> bool:
        """Reject delayed hooks from a prior run that reused the task ID."""
        turn_id = _text(event, "turn_id")
        if not turn_id:
            return True
        return turn_id not in task.retired_turn_ids and (
            not task.turn_ids or turn_id in task.turn_ids
        )

    def _admits(
        self,
        session: _MetricsSession,
        task: _TaskRun,
        event: dict[str, Any],
        *,
        current: bool = False,
    ) -> bool:
        """Whether ``event`` may act on ``task`` (caller holds ``session.lock``).

        Rejects closing sessions and stale turns; with ``current`` also requires ``task`` to
        still be the session's live run for its ID. Admitted events have their turn remembered.
        """
        if (
            session.closing
            or not self._event_matches_task_turn(task, event)
            or (current and session.tasks.get(task.task_id) is not task)
        ):
            return False
        self._remember_turn(session, task, event)
        return True

    def _approval_task(
        self, event: dict[str, Any]
    ) -> tuple[_MetricsSession | None, _TaskRun | None]:
        """Resolve approval correlation without guessing across ambiguous turns."""
        active = relay_runtime.active_turn()
        if active is not None:
            session, task = self._task_pair(
                {**event, "session_id": active.lease.session_id, "task_id": active.task_id}
            )
            if task is not None:
                return session, task

        session, task = self._task_pair(event)
        if task is not None:
            return session, task

        turn_id = _text(event, "turn_id")
        session = None
        if turn_id:
            with self._task_sessions_lock:
                session = _sole(
                    candidate
                    for (owner_id, candidate_turn_id), candidate in self._turn_sessions.items()
                    if candidate_turn_id == turn_id and self._sessions.get(owner_id) is candidate
                )
        if session is None:
            return None, None
        task = _sole(task for task in session.tasks.values() if turn_id in task.turn_ids)
        return (None, None) if task is None else (session, task)

    def _open_tool_call(self, task: _TaskRun, event: dict[str, Any]) -> _ToolCall:
        handle = self._run_in_task(
            task, self.relay.tools.call, contract.TOOL_CALL_SCOPE, {},
            handle=task.handle, metadata=self._event_metadata(),
        )
        return _ToolCall(
            handle, contract.tool_category(event), contract.tool_metric_name(event), monotonic_ns()
        )

    def _finish_tool_call(
        self, task: _TaskRun, tool_call: _ToolCall, event: dict[str, Any]
    ) -> None:
        fields = contract.tool_terminal_fields(
            event, category=tool_call.category, approval_outcome=tool_call.approval_outcome,
            fallback_duration_ms=_elapsed_ms(tool_call.started_ns), tool_name=tool_call.tool_name,
        )
        self._guarded(
            "Hermes shared-metrics tool call close failed",
            lambda: self._run_in_task(
                task, self.relay.tools.call_end, tool_call.handle,
                self.relay.ToolExecutionResult(fields),
                metadata=self._event_metadata(),
            ),
        )

    def _end_pending_tool_calls(
        self, session: _MetricsSession, task: _TaskRun, event: dict[str, Any]
    ) -> None:
        task_outcome, _, _ = contract.task_terminal_state(event)
        status = {"cancelled": "cancelled", "timed_out": "timeout"}.get(task_outcome, "error")
        for key in [key for key in session.tool_calls if key[0] == task.task_id]:
            self._finish_tool_call(task, session.tool_calls.pop(key), {**event, "status": status})

    def _finish_model_call(
        self, session: _MetricsSession, model_call_key: tuple[str, str], outcome: str
    ) -> None:
        model_call = session.model_calls.pop(model_call_key, None)
        if model_call is None:
            return
        error_class = model_call.error_class
        if outcome == "failed" and error_class == "none":
            error_class = "unknown"
        fields = contract.model_route_fields(
            model_call.fields, call_role="primary", outcome=outcome, error_class=error_class,
            ttft_bucket=model_call.ttft_bucket,
        )
        self._guarded(
            "Hermes shared-metrics model call close failed",
            self._run_scoped, session, session.tasks.get(model_call.task_id),
            self.relay.llm.call_end, model_call.handle, fields,
            metadata=self._event_metadata(),
        )

    def _end_pending_model_calls(self, session: _MetricsSession, event: dict[str, Any]) -> None:
        """Close calls that never saw a successful response as failed or cancelled."""
        task_id = _text(event, "task_id")
        cancelled = contract.task_terminal_state(event)[0] == "cancelled"
        pending = [k for k, c in session.model_calls.items() if not task_id or c.task_id == task_id]
        for key in pending:
            self._finish_model_call(session, key, "cancelled" if cancelled else "failed")

    @staticmethod
    def _existing_model_call_key(
        session: _MetricsSession, event: dict[str, Any]
    ) -> tuple[str, str] | None:
        """(task_id, request_id) of an open call; a task-less event may match by request alone."""
        request_id = _text(event, "api_request_id")
        if not request_id:
            return None
        key = (_text(event, "task_id"), request_id)
        if key in session.model_calls or key[0]:
            return key if key in session.model_calls else None
        candidates = [candidate for candidate in session.model_calls if candidate[1] == request_id]
        return candidates[0] if len(candidates) == 1 else None

    def _finish_task(self, session: _MetricsSession, task_id: str, event: dict[str, Any]) -> bool:
        task = session.tasks.get(task_id)
        if task is None:
            return False
        self._end_pending_tool_calls(session, task, event)
        self._end_pending_model_calls(session, {**event, "task_id": task_id})
        fields = contract.task_terminal_fields(
            {**task.start_fields, **event},
            duration_ms=_elapsed_ms(task.started_ns),
            model_call_count=len(task.model_call_ids),
            tool_call_count=len(task.tool_call_ids) + task.unidentified_tool_calls,
            retry_count=task.retry_count,
        )
        if task.cost.user_turn:  # background review forks share the session id and never fire pre_llm_call
            self._count_session_turn(session, task, fields["outcome"])
        if not session.closing:
            self._observe_model_turn(session, task, fields)
        try:
            popped = self._guarded(
                "Hermes shared-metrics task close failed",
                self._run_in_task, task, relay_runtime.pop_relay_scope_if_top, self.relay, task.handle,
                output=fields, metadata=self._event_metadata(),
            )
            if popped is False:
                logger.debug("Left shared-metrics task scope %s under a concurrent turn's scope; session close drains it", task_id)
        finally:
            session.tasks.pop(task_id, None)
            session.retired_turn_ids.extend(task.turn_ids)
            with self._task_sessions_lock:
                _forget(self._task_sessions, (session.session_id, task_id), session)
                for turn_id in task.turn_ids:
                    _forget(self._turn_sessions, (session.session_id, turn_id), session)
        return True

    @staticmethod
    def _count_session_turn(session: _MetricsSession, task: _TaskRun, outcome: str) -> None:
        if session.start_fields is None:
            session.start_fields = dict(task.start_fields)
            session.first_turn_ns = task.started_ns
        session.turns += 1
        session.failed_turns += outcome == "failed"
        session.api_calls += len(task.model_call_ids)
        session.tool_call_total += len(task.tool_call_ids) + task.unidentified_tool_calls
        session.last_outcome = outcome
        session.last_turn_ns = monotonic_ns()

    def _observe_model_turn(self, session: _MetricsSession, task: _TaskRun, fields: dict[str, str]) -> None:
        """A turn the user saw end (session-close aborts excluded): trailing-failure state and interrupts."""
        route = task.model_route or session.model_state.last_route
        session.model_state.observe_turn(fields["outcome"], route, session.last_turn_ns)
        if task.selected_route is not None and task.cost.user_turn:
            self._with_route_run(session.session_id, session.route_run, lambda run: run.observe(task.selected_route))
        if task.model_route is not None and engagement_.engaged_turn(task.start_fields, task.cost.user_turn):
            self._guarded(
                "Hermes shared-metrics engagement mark failed", self._mark,
                session, None, contract.ENGAGEMENT_TURN_MARK, dict(task.model_route),
            )
        if fields["end_reason"] == "user_cancelled" and route is not None and model_.attended(task.start_fields):
            self._guarded(
                "Hermes shared-metrics friction mark failed", self._mark,
                session, None, contract.MODEL_FRICTION_MARK, model_.friction_fields("interrupt", route),
            )
        if task.cost.user_turn and model_.attended(task.start_fields):
            rows = session.efficiency.finish_turn(
                task.cost, fields, route=route, model_calls=len(task.model_call_ids),
                tool_calls=len(task.tool_call_ids) + task.unidentified_tool_calls,
            )
            self._emit_rows(session, rows)

    def _emit_rows(self, session: _MetricsSession | None, rows: list[tuple[str, dict[str, str]]]) -> None:
        for mark, data in rows:
            if session is None:
                self._guarded("Hermes shared-metrics efficiency mark failed", self.record_process_mark, mark, data)
            else:
                self._guarded("Hermes shared-metrics efficiency mark failed", self._mark, session, None, mark, data)

    def _observe_turn_call(
        self, session: _MetricsSession, task: _TaskRun | None, model_call: _ModelCall, event: dict[str, Any]
    ) -> None:
        """One finished primary call of a user turn: its tokens, and a cold cache read nobody announced."""
        if task is None or not task.cost.user_turn:
            return
        route, usage = model_call.fields, event.get("usage")
        task.cost.add_call(route, usage)
        expected = self._consume_cold(self._cold_key(session.session_id))
        cause = session.efficiency.observe_cache(route, usage, model_call.started_ns, monotonic_ns(), expected=expected)
        if cause is not None:
            self._emit_rows(session, [eff.cache_break_row(cause, route)])

    @staticmethod
    def _cold_key(session_id: str) -> str:
        """The conversation (lineage root survives compression rotation), else the session."""
        return get_conversation_context() or session_id

    def _announce_cold(self, key: str) -> bool:
        """Expect the next read of ``key`` to be cold; False when a break was already announced."""
        with self._cold_lock:
            if key in self._cold_expected:
                return False
            self._cold_expected[key] = None
            while len(self._cold_expected) > 256:
                self._cold_expected.popitem(last=False)
            return True

    def _consume_cold(self, key: str) -> bool:
        with self._cold_lock:
            return self._cold_expected.pop(key, False) is None

    def record_known_cache_break(self, cause: str, route: dict[str, str], session_id: str) -> None:
        """Hermes invalidated the prefix itself: count it, and don't recount the cold read it causes.
        Several causes before one cold read are one break (the first cause names it)."""
        if self._announce_cold(self._cold_key(session_id)):
            self._emit_rows(None, [eff.cache_break_row(cause, route)])

    def record_session_tools(self, session_id: str, agent: Any, tools_for_api: list) -> None:
        session = self._session({"session_id": session_id}) if session_id else None
        if session is None:
            return
        key = (id(getattr(agent, "tools", None)), id(tools_for_api), len(tools_for_api))
        with session.lock:
            if session.closing or session.efficiency.tools_key == key:
                return
        snapshot = eff.tool_snapshot(agent, tools_for_api)  # registry + estimator: outside the lock
        with session.lock:
            if session.closing:
                return
            session.efficiency.tools_key = key
            changed = session.efficiency.tools.observe(*snapshot)
        route = model_.model_route(getattr(agent, "provider", None), getattr(agent, "model", None))
        if changed and self._announce_cold(self._cold_key(session_id)):
            self._emit_rows(session, [eff.cache_break_row("toolset_change", route)])

    def rotate_segment(self, old_id: str, new_id: str) -> None:
        """Compression rotated ``old_id`` to ``new_id``: the new id continues the same conversation.
        Only this explicit hand-off links segments (a reset, /new or /branch parent is a new
        conversation). The old segment closes into the lineage now, or once its in-flight turn ends."""
        old = self._session({"session_id": old_id})
        with self._lineage_lock:
            key = self._lineage_of.get(old_id)
            if key is None:
                if old is None:
                    return  # a segment this process never served: the new id is its own conversation
                key = old_id
                self._lineages[key] = _PeakLineage(open={old_id}, run=old.route_run, spent=old.efficiency.spent)
                self._lineage_of[old_id] = key
            self._lineages[key].open.add(new_id)
            self._lineage_of[new_id] = key
            overflow = list(self._lineages)[:-_MAX_LINEAGES]
        new = self._session({"session_id": new_id})
        if new is not None:
            self._share_spent(new)
        if old is not None:
            with old.lock:
                old.retiring = not old.closing
                idle = old.retiring and not old.tasks
            if idle:
                self.close_session({"session_id": old_id})
        for stale in overflow:
            self._flush_lineage(stale)

    def _lineage_spent(self, session_id: str) -> deque | None:
        with self._lineage_lock:
            lineage = self._lineages.get(self._lineage_of.get(session_id, ""))
            return lineage.spent if lineage is not None else None

    def _share_spent(self, session: _MetricsSession) -> None:
        """A rotated-to segment's /undo and /retry reach the turns the earlier segments spent."""
        spent = self._lineage_spent(session.session_id)
        if spent is not None:
            session.efficiency.spent = spent

    def _with_route_run(
        self, session_id: str, own: engagement_.RouteRun | None, update: Callable[[engagement_.RouteRun], Any],
    ) -> Any:
        """Apply ``update`` to the conversation's model-turn run (the lineage's, shared by its segments)."""
        with self._lineage_lock:
            lineage = self._lineages.get(self._lineage_of.get(session_id, ""))
            run = lineage.run if lineage is not None else own
            return update(run) if run is not None else None

    def _close_segment(
        self, session_id: str, state: model_.ModelSessionState | None, tally: engagement_.SessionTally | None,
        tools: eff.ToolUsage,
    ) -> tuple[model_.ModelSessionState | None, engagement_.SessionTally | None, eff.ToolUsage | None]:
        """This segment's context peak, tally and tool usage, or, for a lineage segment, the
        conversation's merged ones once its last open segment closes (None before that)."""
        with self._lineage_lock:
            key = self._lineage_of.pop(session_id, None)
            lineage = self._lineages.get(key) if key is not None else None
            if lineage is None:
                return state, tally, tools
            lineage.open.discard(session_id)
            if state is not None and tally is not None:
                _absorb_peak(lineage.peak, state)
                lineage.tally.absorb(tally)
            lineage.tools.absorb(tools)
            if lineage.open:
                return None, None, None
            del self._lineages[key]
        return lineage.peak, lineage.tally, lineage.tools

    def _close_unseen_segment(self, session_id: str) -> None:
        """A rotated-to segment closed (or the process ended) before it served a turn here."""
        if session_id:
            self._emit_summary_rows(None, *self._close_segment(session_id, None, None, eff.ToolUsage()))

    def _flush_lineage(self, key: str) -> None:
        """Backstop for a conversation a surface never closed: report what it holds, detach its segments."""
        with self._lineage_lock:
            lineage = self._lineages.pop(key, None)
            for segment in lineage.open if lineage is not None else ():
                self._lineage_of.pop(segment, None)
        if lineage is not None:
            self._emit_summary_rows(None, lineage.peak, lineage.tally, lineage.tools)

    def _emit_summary_rows(
        self, session: _MetricsSession | None, peak: model_.ModelSessionState | None,
        tally: engagement_.SessionTally | None, tools: eff.ToolUsage | None, *extra: tuple[str, Any],
    ) -> None:
        marks = [
            (contract.CONTEXT_PEAK_MARK, peak.context_peak_fields() if peak is not None else None),
            (contract.SESSION_MARK, tally.fields() if tally is not None else None),
            *(tools.rows() if tools is not None else ()),
            *extra,
        ]
        self._emit_rows(session, [(mark, data) for mark, data in marks if data is not None])

    def _emit_session_summary(self, session: _MetricsSession) -> None:
        """One hermes.session.count and context-peak row per closed top-level conversation (delegated
        children excluded; compression-rotated segments merged), its tool usage, plus a quick-abandon friction."""
        start = session.start_fields
        counted = bool(session.turns) and start is not None and start.get("entrypoint") != "delegated"
        tally = engagement_.SessionTally(
            session.start_fields, session.turns, session.failed_turns, session.last_outcome,
            session.first_turn_ns, session.last_turn_ns, session.api_calls, session.tool_call_total, session.replies,
        ) if counted else None
        tools = session.efficiency.tools if counted and model_.attended(start) else eff.ToolUsage()
        tools.surface = (start or {}).get("execution_surface", "unknown")
        merged = self._close_segment(session.session_id, session.model_state if counted else None, tally, tools)
        # A segment compression retired is not the user walking away.
        abandoned = session.model_state.quick_abandon_route(monotonic_ns()) if counted and not session.retiring else None
        friction = (
            [(contract.MODEL_FRICTION_MARK, model_.friction_fields("quick_abandon", abandoned))]
            if abandoned is not None and model_.attended(start) else []
        )
        self._emit_summary_rows(session, *merged, *friction)

    def record_switch_after(self, session_id: str) -> None:
        """A counted /model switch: how many turns the model being left served in this conversation."""
        session = self._session({"session_id": session_id})
        with session.lock if session is not None else contextlib.nullcontext():
            taken = self._with_route_run(
                session_id, session.route_run if session is not None else None, lambda run: run.take()
            )
        if taken is not None:
            self.record_process_mark(contract.MODEL_SWITCH_AFTER_MARK, engagement_.switch_after_fields(*taken))

    def record_friction(self, signal: str, session_id: str, fallback_route: dict[str, str], turns: int = 1) -> None:
        """A user friction action, attributed to the session's last primary model when known; an
        /undo or /retry also counts the tokens of the turns it threw away."""
        session = self._session({"session_id": session_id}) if session_id else None
        route, wasted = None, []
        # A session this process never saw (restart, remote host) still counts, with tokens unknown.
        with session.lock if session is not None else contextlib.nullcontext():
            if session is not None:
                route = session.model_state.last_route
            if signal in contract.WASTE_REASONS:
                efficiency = session.efficiency if session is not None else eff.SessionEfficiency()
                if session is None and (spent := self._lineage_spent(session_id)) is not None:
                    efficiency.spent = spent  # a rotated-to segment that has not served a turn yet
                wasted = efficiency.discard_turns(signal, turns, route or fallback_route)
        data = model_.friction_fields(signal, route or fallback_route)
        if data is not None:
            self.record_process_mark(contract.MODEL_FRICTION_MARK, data)
            self._emit_rows(None, wasted)

    def record_process_mark(self, mark: str, data: dict[str, Any]) -> None:
        """A process-level fact with no owning task (setup, installs, commands, switches)."""
        # No flush here: the next task close or the atexit shutdown drains the subscriber.
        self._with_scope_stack(self.relay.scope.event, mark, data=data, metadata=self._event_metadata())

    def record_process_marks_saved(self, marks: list[tuple[str, dict[str, Any]]]) -> int:
        """Emit ``marks`` and wait for the store; how many settled without a store error."""
        ticket = uuid.uuid4().hex
        metadata = {**self._event_metadata(), contract.COMMIT_TICKET_KEY: ticket}
        for mark, data in marks:
            self._with_scope_stack(self.relay.scope.event, mark, data=data, metadata=metadata)
        self.relay.subscribers.flush()
        return self.subscriber.take_saved(ticket)

    def record_auxiliary_tokens(self, event: dict[str, Any]) -> None:
        tokens = fields_.model_token_fields(
            event.get("usage"), call_role="auxiliary", aux_task=event.get("aux_task"),
            model=event.get("response_model") or event.get("model"), provider=event.get("provider"),
        )
        if tokens is not None:
            self._with_scope_stack(
                self.relay.scope.event, contract.MODEL_TOKENS_MARK, data=tokens,
                metadata=self._event_metadata(),
            )

    def _export(self) -> None:
        exported = self._safe(self.subscriber.store.create_and_export_package_if_due)
        # Sending must never delay the caller: _export runs on finish_task, the user's
        # interactive path. The thread is about latency, not correctness.
        if exported is not None:
            self._safe(self._send_exported_packages)

    def _send_exported_packages(self) -> None:
        try:
            resolved = _resolved_send_config()
        except Exception:
            logger.debug("Unable to read shared-metrics send policy", exc_info=True)
            return

        # Observe the consent EDGE before deciding whether to send: the dominant revocation
        # case is "sending turned off while no pass is running", invisible to the send loop.
        # Failures never break the export hook but log at warning (privacy-relevant).
        self._guarded(
            "Unable to record a shared-metrics consent transition",
            _reconcile_store_consent, self.subscriber.store, resolved.send,
        )
        if not resolved.send:
            return

        with self._send_lock:
            # One in-flight pass per process; the next hook fire picks up what is pending.
            if self._send_thread is None or not self._send_thread.is_alive():
                self._send_thread = threading.Thread(
                    target=self._run_send_pass, args=(resolved.endpoint,),
                    name="hermes-shared-metrics-send", daemon=True,
                )
                self._send_thread.start()

    def _run_send_pass(self, endpoint: str) -> None:
        from hermes_cli.observability.shared_metrics_sender import SharedMetricsSender

        def still_consented() -> bool:
            """Re-read consent so revoking `send` stops an in-flight pass."""
            resolved = _resolved_send_config()
            return resolved.send and resolved.endpoint == endpoint

        sender = SharedMetricsSender(self.subscriber.store, endpoint, consent_check=still_consented)
        self._guarded("Shared-metrics send pass failed", sender.send_pending)

    def _event_metadata(self) -> dict[str, str]:
        return {
            contract.SCHEMA_KEY: contract.SCHEMA_VERSION,
            relay_runtime.RUNTIME_INSTANCE_KEY: self.host.runtime_id,
        }

    @staticmethod
    def _guarded(message: str, callback: Callable[..., Any], *args: Any, **kwargs: Any) -> Any:
        """Run ``callback``; log-and-swallow any exception, returning None."""
        try:
            return callback(*args, **kwargs)
        except Exception:
            logger.warning(message, exc_info=True)
            return None

    @classmethod
    def _safe(cls, callback: Callable[..., Any], *args: Any, **kwargs: Any) -> Any:
        return cls._guarded("Hermes shared metrics operation failed", callback, *args, **kwargs)


def _raw_config() -> dict[str, Any]:
    """Read-only config snapshot (lazy import: tests patch ``hermes_cli.config``).

    Collection consent is profile-owned: managed overlays cannot opt a profile in or out.
    The read-only path matters because this gate runs 2-3x per agent turn and the mutable
    read_raw_config() paid a full config deepcopy on every call.
    """
    from hermes_cli.config import read_raw_config_readonly

    return read_raw_config_readonly() or {}


def _resolved_send_config():
    """Resolve the opt-in send policy from the read-only config snapshot."""
    from hermes_cli.observability.shared_metrics_send_config import resolve_send_config

    return resolve_send_config(_raw_config())


def _reconcile_store_consent(store: SharedMetricsStore, send_enabled: bool) -> None:
    from hermes_cli.observability.shared_metrics_sender import reconcile_send_consent
    from hermes_cli.sqlite_util import write_txn

    with store._connection() as connection:
        with write_txn(connection):
            reconcile_send_consent(connection, send_enabled)


def enabled() -> bool:
    """Return the shared-metrics policy for the active Hermes profile."""
    profile_key = relay_runtime.current_profile_key()
    try:
        config: Any = _raw_config()
    except Exception:
        logger.debug("Unable to read Hermes shared-metrics policy", exc_info=True)
        config = None
    for key in ("telemetry", "shared_metrics"):
        config = config.get(key) if isinstance(config, dict) else None
    if isinstance(config, dict) and config.get("enabled") is True:
        return True
    if profile_key not in _RUNTIMES:  # the opted-out hot path: nothing to tear down, no lock
        return False
    with _RUNTIME_LOCK:
        runtime = _RUNTIMES.pop(profile_key, None)
        if isinstance(runtime, _Runtime):
            runtime.deactivate()
    return False


def handles_hook(hook_name: str) -> bool:
    return hook_name in HANDLED_HOOKS and enabled()


_consent_reconcile_done = False


def _reconcile_send_consent_once() -> None:
    """Reconcile consent windows with config, once per process.

    Runs BEFORE and INDEPENDENT of the collection gate, so a user with ``enabled: false``
    still gets send-consent windows reconciled. Skipped only when there is no store on disk
    AND consent is off: nothing to protect, and creating ``~/.hermes/telemetry`` for every
    fully-disabled user would be the wrong behaviour change.
    """
    global _consent_reconcile_done
    if _consent_reconcile_done:
        return
    _consent_reconcile_done = True
    try:
        # Lazy: tests patch ``shared_metrics.SharedMetricsStore`` at its origin.
        from hermes_cli.observability.shared_metrics import SharedMetricsStore
        from hermes_constants import get_hermes_home

        resolved = _resolved_send_config()
        # Probe WITHOUT constructing a store: the constructor creates the directory and
        # schema as a side effect, which would make the skip below dead code.
        default_path = get_hermes_home() / "telemetry" / "shared_metrics" / "metrics.sqlite3"
        if not resolved.send and not default_path.exists():
            return
        _reconcile_store_consent(SharedMetricsStore(), resolved.send)
    except Exception:
        logger.warning("Unable to reconcile shared-metrics send consent", exc_info=True)


def observe_lifecycle(hook_name: str, **kwargs: Any) -> None:
    """Project one Hermes lifecycle event into the core Relay integration."""
    _reconcile_send_consent_once()
    if not handles_hook(hook_name) or not relay_runtime.relay_instrumentation_enabled():
        return
    runtime = _get_runtime()
    if runtime is None:
        return
    try:
        _HOOK_HANDLERS[hook_name](runtime, kwargs)
    except Exception:
        logger.warning("Hermes shared metrics hook failed: %s", hook_name, exc_info=True)


def _with_runtime_toolset(event: dict[str, Any]) -> dict[str, Any]:
    """Attach the toolset already declared by Hermes's runtime registry."""
    tool_name = _text(event, "tool_name")
    if event.get("toolset") or not tool_name:
        return event
    try:
        from model_tools import get_toolset_for_tool

        toolset = get_toolset_for_tool(tool_name)
    except Exception:
        toolset = None
    return {**event, "toolset": toolset or "other"}


def _close_child_session(runtime: _Runtime, kwargs: dict[str, Any]) -> None:
    child_session_id = _text(kwargs, "child_session_id")
    if child_session_id:
        runtime.close_session({"session_id": child_session_id})


_HOOK_HANDLERS: dict[str, Callable[[_Runtime, dict[str, Any]], Any]] = {
    "on_session_start": lambda rt, kw: rt.record_client_active(kw),
    "pre_llm_call": lambda rt, kw: rt.start_user_turn(kw),
    "pre_api_request": lambda rt, kw: rt.start_model_call(kw),
    "pre_tool_call": lambda rt, kw: rt.start_tool_call(_with_runtime_toolset(kw)),
    "post_tool_call": lambda rt, kw: rt.record_tool_call(_with_runtime_toolset(kw)),
    "post_approval_response": lambda rt, kw: rt.record_approval(kw),
    "on_skill_lifecycle": lambda rt, kw: rt.record_skill_lifecycle(kw),
    "post_api_request": lambda rt, kw: rt.update_model_call(kw, finish=True),
    "post_auxiliary_call": lambda rt, kw: rt.record_auxiliary_tokens(kw),
    "api_request_error": lambda rt, kw: rt.update_model_call(kw, finish=False),
    "on_session_end": lambda rt, kw: rt.finish_task(kw),
    "subagent_stop": _close_child_session,
    "on_session_finalize": lambda rt, kw: rt.close_session(kw),
    "on_session_reset": lambda rt, kw: rt.close_session(kw),
}
HANDLED_HOOKS = frozenset(_HOOK_HANDLERS)


def _prepare_core_session(host: relay_runtime.RelayRuntime, context: dict[str, Any]) -> None:
    """Prepare the profile subscriber before the coordinator opens a scope."""
    del context
    if host.profile_key == relay_runtime.current_profile_key() and enabled():
        _get_runtime(retry_failed=True, host=host)


def start_task_run(
    *, session_id: str, task_id: str, platform: str, parent_session_id: str = ""
) -> None:
    """Start task metrics at the outer Hermes execution boundary."""
    _run_task_hook(
        "start_task", retry_failed=True, session_id=session_id, task_id=task_id,
        platform=platform, parent_session_id=parent_session_id,
    )


def close_session_run(session_id: str) -> None:
    """Emit the summary of a session a surface retired without a finalize hook (gateway
    auto-reset, idle expiry). Never creates a runtime; unknown ids are a no-op."""
    if not session_id or not enabled():
        return
    runtime = _RUNTIMES.get(relay_runtime.current_profile_key())
    if isinstance(runtime, _Runtime):
        runtime._safe(runtime.close_session, {"session_id": session_id})


def rotate_segment(old_session_id: str, new_session_id: str) -> None:
    """Compression rotated a conversation's session id (the one seam that continues a conversation
    under a new id). Never creates a runtime: a runtime that does not exist served no segment."""
    if not old_session_id or not new_session_id or old_session_id == new_session_id or not enabled():
        return
    runtime = _RUNTIMES.get(relay_runtime.current_profile_key())
    if isinstance(runtime, _Runtime):
        runtime._safe(runtime.rotate_segment, old_session_id, new_session_id)


def finish_task_run(
    *, session_id: str, task_id: str, platform: str,
    result: dict[str, Any] | None = None, error: BaseException | None = None,
) -> None:
    """Finish task metrics for every return or exception path."""
    _run_task_hook(
        "finish_task", session_id=session_id, task_id=task_id, platform=platform,
        **_terminal_flags(result, error),
    )


def record_process_mark(mark: str, data: dict[str, Any]) -> None:
    """Emit one process-level decision mark when shared metrics are on (never raises)."""
    if not enabled() or not relay_runtime.relay_instrumentation_enabled():
        return
    runtime = _get_runtime(retry_failed=True)
    if runtime is not None:
        runtime._safe(runtime.record_process_mark, mark, data)


def record_process_marks_saved(marks: list[tuple[str, dict[str, Any]]]) -> int:
    """``record_process_mark`` for facts recovered from a file the caller deletes only once they are
    saved (a busy store would otherwise lose them for good). How many rows are settled: saved, rejected
    by the contract, or nothing to record because collection is off."""
    if not enabled() or not relay_runtime.relay_instrumentation_enabled():
        return len(marks)
    runtime = _get_runtime(retry_failed=True)
    return (runtime._safe(runtime.record_process_marks_saved, marks) or 0) if runtime is not None else 0


def record_session_friction(signal: str, session_id: str, fallback_route: dict[str, str], turns: int = 1) -> None:
    """Record one friction signal in the active profile's runtime (callers already checked enabled())."""
    if not relay_runtime.relay_instrumentation_enabled():
        return
    runtime = _get_runtime(retry_failed=True)
    if runtime is not None:
        runtime._safe(runtime.record_friction, signal, session_id, fallback_route, turns)


def record_session_tools(session_id: str, agent: Any, tools_for_api: list) -> None:
    """The tools one primary request sent (caller checked enabled()); never creates a runtime."""
    runtime = _RUNTIMES.get(relay_runtime.current_profile_key())
    if isinstance(runtime, _Runtime):
        runtime._safe(runtime.record_session_tools, session_id, agent, tools_for_api)


def record_known_cache_break(cause: str, route: dict[str, str], session_id: str) -> None:
    """A prompt-cache break Hermes caused (caller checked enabled())."""
    if not relay_runtime.relay_instrumentation_enabled():
        return
    runtime = _get_runtime(retry_failed=True)
    if runtime is not None:
        runtime._safe(runtime.record_known_cache_break, cause, route, session_id)


def record_model_switch_after(session_id: str) -> None:
    """Count turns-before-switch for a session this process served; never creates a runtime (a
    runtime that does not exist saw no turns). Callers are already in the owning profile's scope."""
    if not session_id or not enabled() or not relay_runtime.relay_instrumentation_enabled():
        return
    runtime = _RUNTIMES.get(relay_runtime.current_profile_key())
    if isinstance(runtime, _Runtime):
        runtime._safe(runtime.record_switch_after, session_id)


def _run_task_hook(method: str, *, retry_failed: bool = False, **event: Any) -> None:
    if not enabled():
        return
    runtime = _get_runtime(retry_failed=retry_failed)
    if runtime is not None:
        runtime._safe(getattr(runtime, method), event)


def _terminal_flags(result: dict[str, Any] | None, error: BaseException | None) -> dict[str, Any]:
    """Bounded completed/failed/interrupted/turn_exit_reason for a task's return or raise."""
    if error is not None:
        interrupted = (
            isinstance(error, (KeyboardInterrupt, InterruptedError))
            or type(error).__name__ == "CancelledError"
        )
        if interrupted:
            reason = "interrupted_by_user"
        else:
            reason = "timed_out" if isinstance(error, TimeoutError) else "system_aborted"
        return {
            "completed": False, "failed": not interrupted, "interrupted": interrupted,
            "turn_exit_reason": reason, "failure_class": "exception",
        }
    terminal = result if isinstance(result, dict) else {}
    failed = terminal.get("failed") is True
    reason = str(terminal.get("turn_exit_reason") or terminal.get("failure_reason") or "")
    return {
        "completed": terminal.get("completed") is True,
        "failed": failed,
        "interrupted": terminal.get("interrupted") is True,
        "turn_exit_reason": reason or ("failed" if failed else "unknown"),
        # Classified in the contract against a closed set; the raw string never leaves.
        "failure_reason": str(terminal.get("failure_reason") or ""),
    }


def _get_runtime(
    *, retry_failed: bool = False, host: relay_runtime.RelayRuntime | None = None
) -> _Runtime | None:
    profile_key = relay_runtime.current_profile_key()
    with _RUNTIME_LOCK:
        runtime = _RUNTIMES.get(profile_key)
        if isinstance(runtime, _Runtime):
            if host is None or runtime.host is host:
                return runtime
            runtime.deactivate()
        elif runtime is _RUNTIME_FAILED and not retry_failed:
            return None
        try:
            _RUNTIMES[profile_key] = runtime = _Runtime(host=host)
        except Exception:
            logger.warning("Hermes shared metrics initialization failed", exc_info=True)
            _RUNTIMES[profile_key] = _RUNTIME_FAILED
            return None
        return runtime


relay_runtime.SESSION_COORDINATOR.register_session_initializer(
    SUBSCRIBER_NAME, _prepare_core_session
)


def _reset_for_tests() -> None:
    """Reset all profile-scoped shared-metrics state for isolated tests."""
    with _RUNTIME_LOCK:
        runtimes = list(_RUNTIMES.values())
        _RUNTIMES.clear()
    for runtime in runtimes:
        if isinstance(runtime, _Runtime):
            runtime.shutdown()
