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

from __future__ import annotations

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

from agent import relay_runtime
from hermes_cli import __version__

from .shared_metrics import SharedMetricsStore
from .shared_metrics_contract import (
    CLIENT_ACTIVE_MARK,
    MODEL_CALL_PROFILE_MODEL,
    MODEL_CALL_SCOPE,
    SCHEMA_KEY,
    SCHEMA_VERSION,
    SKILL_LIFECYCLE_MARK,
    SKILL_LOAD_MARK,
    SUBSCRIBER_NAME,
    TASK_SCOPE,
    TOOL_APPROVAL_MARK,
    TOOL_CALL_SCOPE,
    model_call_fields,
    skill_lifecycle_fields,
    skill_load_fields,
    task_start_fields,
    task_terminal_fields,
    task_terminal_state,
    tool_approval_outcome,
    tool_category,
    tool_terminal_fields,
)
from .shared_metrics_subscriber import SharedMetricsSubscriber

logger = logging.getLogger(__name__)

HANDLED_HOOKS = frozenset({
    "on_session_start",
    "on_session_end",
    "on_session_finalize",
    "on_session_reset",
    "pre_llm_call",
    "pre_api_request",
    "pre_tool_call",
    "post_tool_call",
    "post_approval_response",
    "post_api_request",
    "api_request_error",
    "on_skill_lifecycle",
    "subagent_stop",
})

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


def _retry_ordinal(event: dict[str, Any]) -> int | None:
    value = event.get("retry_count")
    if isinstance(value, int) and not isinstance(value, bool) and value >= 0:
        return value
    return None


@dataclass
class _ModelCall:
    handle: Any
    task_id: str
    fields: dict[str, str]
    retry_ordinal: int | None = None


@dataclass
class _ToolCall:
    handle: Any
    task_id: str
    category: 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


@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),
    )


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._sessions_lock = threading.RLock()
        self._active = True
        self._sessions: dict[str, _MetricsSession] = {}
        self._task_creation_lock = threading.RLock()
        self._task_sessions_lock = threading.RLock()
        self._task_sessions: dict[tuple[str, str], _MetricsSession] = {}
        self._turn_sessions: dict[tuple[str, str], _MetricsSession] = {}
        self._subscriber_name = f"{SUBSCRIBER_NAME}.{self.host.runtime_id}"
        self.subscriber = SharedMetricsSubscriber(
            SharedMetricsStore(),
            __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 = str(event.get("session_id") or "")
        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
        with session.lock:
            if session.closing:
                return None
        return 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 None:
            return
        self._emit_client_active(session)

    def _emit_client_active(self, session: _MetricsSession) -> None:
        with session.lock:
            if session.closing:
                return
            self._run_in_session(
                session,
                self.relay.scope.event,
                CLIENT_ACTIVE_MARK,
                handle=session.relay_session.handle,
                data={},
                metadata=self._event_metadata(),
            )

    def _run_in_session(
        self,
        session: _MetricsSession,
        callback: Callable[..., Any],
        *args: Any,
        **kwargs: Any,
    ) -> Any:
        return self.host.run_in_session(
            session.relay_session,
            callback,
            *args,
            **kwargs,
        )

    def start_task(self, event: dict[str, Any]) -> _TaskRun | None:
        """Open one Relay function scope for a Hermes task run."""
        task_key = self._task_key(event)
        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:
                        if not self._event_matches_task_turn(task, event):
                            return None
                        self._remember_turn(owner, task, event)
                    return task

            session = self.ensure_session(event)
            if session is None:
                return None
            with session.lock:
                turn_id = str(event.get("turn_id") or "")
                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 = task_start_fields(event)
                active_turn = relay_runtime.active_turn(session.session_id)
                parent_handle = session.relay_session.handle
                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
                ):
                    parent_handle = active_turn.handle

                def push_task() -> Any:
                    self.relay.get_scope_stack()
                    return self.relay.scope.push(
                        TASK_SCOPE,
                        self.relay.ScopeType.Function,
                        handle=parent_handle,
                        input=start_fields,
                        metadata=self._event_metadata(),
                    )

                handle = task_context.run(push_task)
                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 _run_in_task(
        self,
        task: _TaskRun,
        callback: Callable[..., Any],
        *args: Any,
        **kwargs: Any,
    ) -> Any:
        def invoke() -> Any:
            self.relay.get_scope_stack()
            return callback(*args, **kwargs)

        return task.context.copy().run(invoke)

    def start_model_call(self, event: dict[str, Any]) -> None:
        task_id = str(event.get("task_id") or "")
        session = self._task_session(event, allow_task_id_fallback=True)
        task = session.tasks.get(task_id) if session is not None else None
        if task is None:
            task = self.start_task(event)
            session = self._task_session(event) if task is not None else None
            if task_id and task is None:
                return
        if session is None:
            session = self.ensure_session(event)
        if session is None:
            return
        model_call_key = self._new_model_call_key(event)
        if model_call_key is None:
            return
        _, request_id = model_call_key
        fields = model_call_fields(event)
        retry_ordinal = _retry_ordinal(event)
        with session.lock:
            if session.closing:
                return
            if task is not None:
                if (
                    session.tasks.get(task.task_id) is not task
                    or not self._event_matches_task_turn(task, event)
                ):
                    return
                self._remember_turn(session, task, event)
            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 retry counter.
                    task.retry_count += 1
                if retry_ordinal is not None:
                    existing.retry_ordinal = max(
                        existing.retry_ordinal or 0,
                        retry_ordinal,
                    )
                return
            if task is not None:
                task.model_call_ids.add(request_id)
                if retry_ordinal is not None and retry_ordinal > 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_in_task(
                    task,
                    self.relay.llm.call,
                    MODEL_CALL_SCOPE,
                    self.relay.LLMRequest({}, {}),
                    handle=task.handle,
                    metadata=self._event_metadata(),
                    model_name=MODEL_CALL_PROFILE_MODEL,
                )
            else:
                handle = self._run_in_session(
                    session,
                    self.relay.llm.call,
                    MODEL_CALL_SCOPE,
                    self.relay.LLMRequest({}, {}),
                    handle=session.relay_session.handle,
                    metadata=self._event_metadata(),
                    model_name=MODEL_CALL_PROFILE_MODEL,
                )
            session.model_calls[model_call_key] = _ModelCall(
                handle=handle,
                task_id=str(event.get("task_id") or ""),
                fields=fields,
                retry_ordinal=retry_ordinal,
            )

    def record_model_call_error(self, event: dict[str, Any]) -> None:
        """Retain the latest attempt error without closing the logical call."""
        session = self._task_session(event, allow_task_id_fallback=True)
        if session is None:
            session = self._session(event)
        if session is None:
            return
        with session.lock:
            if session.closing:
                return
            model_call_key = self._existing_model_call_key(session, event)
            if model_call_key is None:
                return
            model_call = session.model_calls.get(model_call_key)
            if model_call is None:
                return
            model_call.fields = model_call_fields(event)

    def start_tool_call(self, event: dict[str, Any]) -> None:
        """Open one privacy-safe Relay tool lifecycle under its task."""
        task_id = str(event.get("task_id") or "")
        session = self._task_session(event, allow_task_id_fallback=True)
        task = session.tasks.get(task_id) if session is not None else None
        if task is None:
            task = self.start_task(event)
            session = self._task_session(event) if task is not None else None
        if session is None or task is None:
            return
        tool_call_id = str(event.get("tool_call_id") or "")
        if not tool_call_id:
            return
        identity = self._tool_call_identity(event)
        with session.lock:
            if session.closing:
                return
            if not self._event_matches_task_turn(task, event):
                return
            self._remember_turn(session, task, event)
            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 = tool_approval_outcome(event)
        tool_call_id = str(event.get("tool_call_id") or "")
        attribution = "unattributed"
        with session.lock:
            if session.closing:
                return
            if not self._event_matches_task_turn(task, event):
                return
            if tool_call_id:
                identity = self._tool_call_identity(event)
                tool_call = session.tool_calls.get((task.task_id, *identity))
                if tool_call is None:
                    matching_keys = [
                        key
                        for key in session.tool_calls
                        if key[0] == task.task_id
                        and self._tool_call_identities_are_compatible(
                            key[1:],
                            identity,
                        )
                    ]
                    tool_call = (
                        session.tool_calls[matching_keys[0]]
                        if len(matching_keys) == 1
                        else None
                    )
                if tool_call is not None:
                    tool_call.approval_outcome = outcome
                    attribution = "tool_call"
            self._run_in_task(
                task,
                self.relay.scope.event,
                TOOL_APPROVAL_MARK,
                handle=task.handle,
                data={"attribution": attribution, "outcome": outcome},
                metadata=self._event_metadata(),
            )

    def record_tool_call(self, event: dict[str, Any]) -> None:
        """Close and count one unique privacy-safe tool lifecycle."""
        task_id = str(event.get("task_id") or "")
        session = self._task_session(event, allow_task_id_fallback=True)
        task = session.tasks.get(task_id) if session is not None else None
        if session is None or task is None:
            return
        tool_call_id = str(event.get("tool_call_id") or "")
        with session.lock:
            if session.closing:
                return
            if not self._event_matches_task_turn(task, event):
                return
            self._remember_turn(session, task, event)
            if 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(
                        self._tool_call_identities_are_compatible(
                            completed_identity,
                            observed_identity,
                        )
                        for completed_identity in task.completed_tool_call_ids
                    ):
                        return
                    matching_keys = [
                        key
                        for key in session.tool_calls
                        if key[0] == task_id
                        and self._tool_call_identities_are_compatible(
                            key[1:],
                            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:
                        key = matching_keys[0]
                        identity = key[1:]
                        tool_call = session.tool_calls.pop(key)
                task.completed_tool_call_ids.update({
                    identity,
                    observed_identity,
                })
                task.tool_call_ids.add(identity)
            else:
                task.unidentified_tool_calls += 1
                tool_call = None
            if tool_call is None:
                tool_call = self._open_tool_call(task, event)
            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."""
        action = str(event.get("action") or "").strip().lower()
        if action == "loaded":
            mark = SKILL_LOAD_MARK
            fields = skill_load_fields(event)
        else:
            mark = SKILL_LIFECYCLE_MARK
            fields = skill_lifecycle_fields(event)
        if fields is None:
            return

        session_id = str(event.get("session_id") or "")
        task_id = str(event.get("task_id") or "")
        session = self._task_session(
            event,
            allow_task_id_fallback=not session_id,
        )
        task = session.tasks.get(task_id) if session is not None else None
        if session is not None:
            if task is None:
                return
            with session.lock:
                if session.closing:
                    return
                if (
                    session.tasks.get(task.task_id) is not task
                    or not self._event_matches_task_turn(task, event)
                ):
                    return
                self._run_in_task(
                    task,
                    self.relay.scope.event,
                    mark,
                    handle=task.handle,
                    data=fields,
                    metadata=self._event_metadata(),
                )
            return
        if session_id and task_id:
            return

        self.relay.get_scope_stack()
        self.relay.scope.event(
            mark,
            data=fields,
            metadata=self._event_metadata(),
        )

    def end_model_call(self, event: dict[str, Any]) -> None:
        session = self._task_session(event, allow_task_id_fallback=True)
        if session is None:
            session = self._session(event)
        if session is None:
            return
        with session.lock:
            if session.closing:
                return
            model_call_key = self._existing_model_call_key(session, event)
            if model_call_key is None:
                return
            model_call = session.model_calls.get(model_call_key)
            if model_call is None:
                return
            fields = model_call_fields(event)
            model_call.fields = fields
            self._finish_model_call(
                session,
                model_call_key,
            )

    def end_pending_model_calls(self, event: dict[str, Any]) -> None:
        session = self._task_session(event, allow_task_id_fallback=True)
        if session is None:
            session = self._session(event)
        if session is None:
            return
        with session.lock:
            if session.closing:
                return
            self._end_pending_model_calls(session, event)

    def finish_task(self, event: dict[str, Any]) -> None:
        """Close one task scope exactly once with bounded terminal fields."""
        task_id = str(event.get("task_id") or "")
        session = self._task_session(
            event,
            allow_task_id_fallback=True,
        ) or self._session(event)
        if session is None:
            return
        with session.lock:
            if session.closing:
                return
            finished = self._finish_task(session, task_id, event)
        if finished:
            try:
                self.relay.subscribers.flush()
            except Exception:
                logger.warning(
                    "Hermes shared-metrics task flush failed",
                    exc_info=True,
                )
            else:
                self._export()

    def close_session(self, event: dict[str, Any]) -> None:
        session = self._session(event)
        if session is None:
            return
        failures: list[str] = []
        with session.lock:
            if session.closing:
                return
            session.closing = True
            for task_id in list(session.tasks):
                self._finish_task(
                    session,
                    task_id,
                    {
                        **event,
                        "task_id": task_id,
                        "completed": False,
                        "failed": True,
                        "interrupted": False,
                        "turn_exit_reason": "system_aborted",
                    },
                )
            self._end_pending_model_calls(session, event)
        try:
            self.relay.subscribers.flush()
        except Exception as exc:
            failures.append(f"subscriber flush failed: {exc}")
        else:
            self._export()
        with self._sessions_lock:
            if self._sessions.get(session.session_id) is session:
                self._sessions.pop(session.session_id, None)
        if failures:
            logger.warning(
                "Hermes shared-metrics session %s closed with errors: %s",
                session.session_id,
                "; ".join(failures),
            )

    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})
        if not self._registered:
            return
        try:
            self.relay.subscribers.flush()
        except Exception:
            logger.warning(
                "Hermes shared-metrics shutdown flush failed",
                exc_info=True,
            )
        else:
            self._export()
        self._safe(self.relay.subscribers.deregister, self._subscriber_name)
        self.host.release_managed_execution(self._subscriber_name)
        self._registered = False
        try:
            atexit.unregister(self.shutdown)
        except Exception:
            pass

    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._safe(self.relay.subscribers.deregister, self._subscriber_name)
            self.host.release_managed_execution(self._subscriber_name)
            self._registered = False
        with self._sessions_lock:
            sessions = list(self._sessions.values())
        for session in sessions:
            with session.lock:
                if session.closing:
                    continue
                session.closing = True
                for task_id in list(session.tasks):
                    self._finish_task(
                        session,
                        task_id,
                        {
                            "session_id": session.session_id,
                            "task_id": task_id,
                            "failed": True,
                            "turn_exit_reason": "system_aborted",
                        },
                    )
                self._end_pending_model_calls(session, {})
        with self._sessions_lock:
            self._sessions.clear()
        with self._task_sessions_lock:
            self._task_sessions.clear()
            self._turn_sessions.clear()
        try:
            atexit.unregister(self.shutdown)
        except Exception:
            pass

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

    @staticmethod
    def _task_key(event: dict[str, Any]) -> tuple[str, str] | None:
        session_id = str(event.get("session_id") or "")
        task_id = str(event.get("task_id") or "")
        if not session_id or not task_id:
            return None
        return session_id, task_id

    def _task_session(
        self,
        event: dict[str, Any],
        *,
        allow_task_id_fallback: bool = False,
    ) -> _MetricsSession | None:
        session_id = str(event.get("session_id") or "")
        task_id = str(event.get("task_id") or "")
        if not task_id:
            return None
        task_key = (session_id, task_id) if session_id else None
        turn_key = self._turn_key(event)
        with self._task_sessions_lock:
            if turn_key is not None:
                owner = self._turn_sessions.get(turn_key)
                if owner is not None:
                    return owner
            if task_key is not None:
                owner = self._task_sessions.get(task_key)
                if owner is not None:
                    return owner
            if not allow_task_id_fallback:
                return None
            candidates: list[_MetricsSession] = []
            for (_, candidate_task_id), session in self._task_sessions.items():
                if candidate_task_id != task_id:
                    continue
                if not any(candidate is session for candidate in candidates):
                    candidates.append(session)
            return candidates[0] if len(candidates) == 1 else None

    @staticmethod
    def _turn_key(event: dict[str, Any]) -> tuple[str, str] | None:
        session_id = str(event.get("session_id") or "")
        turn_id = str(event.get("turn_id") or "")
        if not session_id or not turn_id:
            return None
        return session_id, turn_id

    def _remember_turn(
        self,
        session: _MetricsSession,
        task: _TaskRun,
        event: dict[str, Any],
    ) -> None:
        turn_id = str(event.get("turn_id") or "")
        if not turn_id:
            return
        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 (
            str(event.get("api_request_id") or ""),
            str(event.get("turn_id") or ""),
            str(event.get("tool_call_id") or ""),
        )

    @staticmethod
    def _tool_call_identities_are_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 candidate_value
            or not observed_value
            or candidate_value == observed_value
            for candidate_value, observed_value in zip(
                candidate[:2],
                observed[:2],
                strict=True,
            )
        )

    @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 = str(event.get("turn_id") or "")
        if not turn_id:
            return True
        if turn_id in task.retired_turn_ids:
            return False
        return not task.turn_ids or turn_id in task.turn_ids

    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:
            correlated = {
                **event,
                "session_id": active.lease.session_id,
                "task_id": active.task_id,
            }
            session = self._task_session(correlated)
            task = session.tasks.get(active.task_id) if session is not None else None
            if task is not None:
                return session, task

        session = self._task_session(event)
        task_id = str(event.get("task_id") or "")
        task = session.tasks.get(task_id) if session is not None else None
        if task is not None:
            return session, task

        turn_id = str(event.get("turn_id") or "")
        if not turn_id:
            return None, None
        with self._task_sessions_lock:
            candidates = [
                candidate
                for (
                    candidate_session_id,
                    candidate_turn_id,
                ), candidate in self._turn_sessions.items()
                if candidate_turn_id == turn_id
                and self._sessions.get(candidate_session_id) is candidate
            ]
        unique_sessions = {id(candidate): candidate for candidate in candidates}
        if len(unique_sessions) != 1:
            return None, None
        session = next(iter(unique_sessions.values()))
        matching_tasks = [
            candidate
            for candidate in session.tasks.values()
            if turn_id in candidate.turn_ids
        ]
        if len(matching_tasks) != 1:
            return None, None
        return session, matching_tasks[0]

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

    def _finish_tool_call(
        self,
        task: _TaskRun,
        tool_call: _ToolCall,
        event: dict[str, Any],
    ) -> None:
        fields = tool_terminal_fields(
            event,
            category=tool_call.category,
            approval_outcome=tool_call.approval_outcome,
            fallback_duration_ms=max(
                0,
                (monotonic_ns() - tool_call.started_ns) // 1_000_000,
            ),
        )
        try:
            self._run_in_task(
                task,
                self.relay.tools.call_end,
                tool_call.handle,
                fields,
                metadata=self._event_metadata(),
            )
        except Exception:
            logger.warning(
                "Hermes shared-metrics tool call close failed",
                exc_info=True,
            )

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

    def _finish_model_call(
        self,
        session: _MetricsSession,
        model_call_key: tuple[str, str],
    ) -> None:
        model_call = session.model_calls.pop(model_call_key, None)
        if model_call is None:
            return
        try:
            task = session.tasks.get(model_call.task_id)
            if task is not None:
                self._run_in_task(
                    task,
                    self.relay.llm.call_end,
                    model_call.handle,
                    model_call.fields,
                    metadata=self._event_metadata(),
                )
            else:
                self._run_in_session(
                    session,
                    self.relay.llm.call_end,
                    model_call.handle,
                    model_call.fields,
                    metadata=self._event_metadata(),
                )
        except Exception:
            logger.warning(
                "Hermes shared-metrics model call close failed", exc_info=True
            )

    def _end_pending_model_calls(
        self,
        session: _MetricsSession,
        event: dict[str, Any],
    ) -> None:
        task_id = str(event.get("task_id") or "")
        model_call_keys = [
            model_call_key
            for model_call_key, model_call in session.model_calls.items()
            if not task_id or model_call.task_id == task_id
        ]
        for model_call_key in model_call_keys:
            self._finish_model_call(
                session,
                model_call_key,
            )

    @staticmethod
    def _new_model_call_key(event: dict[str, Any]) -> tuple[str, str] | None:
        request_id = str(event.get("api_request_id") or "")
        if not request_id:
            return None
        return str(event.get("task_id") or ""), request_id

    @classmethod
    def _existing_model_call_key(
        cls,
        session: _MetricsSession,
        event: dict[str, Any],
    ) -> tuple[str, str] | None:
        key = cls._new_model_call_key(event)
        if key is None:
            return None
        if key in session.model_calls:
            return key
        if key[0]:
            return None
        candidates = [
            candidate for candidate in session.model_calls if candidate[1] == key[1]
        ]
        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 = task_terminal_fields(
            {**task.start_fields, **event},
            duration_ms=max(0, (monotonic_ns() - task.started_ns) // 1_000_000),
            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,
        )
        try:
            self._run_in_task(
                task,
                self.relay.scope.pop,
                task.handle,
                output=fields,
                metadata=self._event_metadata(),
            )
        except Exception:
            logger.warning("Hermes shared-metrics task close failed", exc_info=True)
        finally:
            session.tasks.pop(task_id, None)
            session.retired_turn_ids.extend(task.turn_ids)
            with self._task_sessions_lock:
                task_key = (session.session_id, task_id)
                if self._task_sessions.get(task_key) is session:
                    self._task_sessions.pop(task_key, None)
                for turn_id in task.turn_ids:
                    turn_key = (session.session_id, turn_id)
                    if self._turn_sessions.get(turn_key) is session:
                        self._turn_sessions.pop(turn_key, None)
        return True

    def _export(self) -> None:
        self._safe(self.subscriber.store.create_and_export_package_if_due)

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

    @staticmethod
    def _safe(callback: Callable[..., Any], *args: Any, **kwargs: Any) -> Any:
        try:
            return callback(*args, **kwargs)
        except Exception:
            logger.warning("Hermes shared metrics operation failed", exc_info=True)
            return None


def enabled() -> bool:
    """Return the shared-metrics policy for the active Hermes profile."""
    profile_key = relay_runtime.current_profile_key()
    try:
        from hermes_cli.config import read_raw_config_readonly

        # Collection consent is profile-owned. Managed config overlays may
        # control runtime policy, but cannot opt a profile into or out of
        # shared metrics. Read-only fast path: this gate runs 2-3x per agent
        # turn, and the mutable read_raw_config() paid a full config deepcopy
        # on every call.
        config = read_raw_config_readonly() or {}
    except Exception:
        logger.debug("Unable to read Hermes shared-metrics policy", exc_info=True)
        value = False
    else:
        telemetry = config.get("telemetry") if isinstance(config, dict) else None
        shared_metrics = (
            telemetry.get("shared_metrics") if isinstance(telemetry, dict) else None
        )
        value = (
            isinstance(shared_metrics, dict) and shared_metrics.get("enabled") is True
        )
    if value:
        return True
    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()


def observe_lifecycle(hook_name: str, **kwargs: Any) -> None:
    """Project one Hermes lifecycle event into the core Relay integration."""
    if not handles_hook(hook_name):
        return
    if not relay_runtime.relay_instrumentation_enabled():
        return
    runtime = _get_runtime()
    if runtime is None:
        return
    try:
        if hook_name == "on_session_start":
            runtime.record_client_active(kwargs)
        elif hook_name == "pre_llm_call":
            runtime.start_task(kwargs)
        elif hook_name == "pre_api_request":
            runtime.start_model_call(kwargs)
        elif hook_name == "pre_tool_call":
            runtime.start_tool_call(_with_runtime_toolset(kwargs))
        elif hook_name == "post_tool_call":
            runtime.record_tool_call(_with_runtime_toolset(kwargs))
        elif hook_name == "post_approval_response":
            runtime.record_approval(kwargs)
        elif hook_name == "on_skill_lifecycle":
            runtime.record_skill_lifecycle(kwargs)
        elif hook_name == "post_api_request":
            runtime.end_model_call(kwargs)
        elif hook_name == "api_request_error":
            runtime.record_model_call_error(kwargs)
        elif hook_name == "on_session_end":
            runtime.finish_task(kwargs)
        elif hook_name == "subagent_stop":
            child_session_id = str(kwargs.get("child_session_id") or "")
            if child_session_id:
                runtime.close_session({"session_id": child_session_id})
        elif hook_name in {"on_session_finalize", "on_session_reset"}:
            runtime.close_session(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."""
    if event.get("toolset"):
        return event
    tool_name = str(event.get("tool_name") or "")
    if 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 prepare_session_start() -> None:
    """Register the subscriber before any producer opens the session scope."""
    if enabled():
        _get_runtime(retry_failed=True)


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():
        if 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."""
    if not enabled():
        return
    runtime = _get_runtime(retry_failed=True)
    if runtime is None:
        return
    runtime._safe(
        runtime.start_task,
        {
            "session_id": session_id,
            "task_id": task_id,
            "platform": platform,
            "parent_session_id": parent_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."""
    if not enabled():
        return
    runtime = _get_runtime()
    if runtime is None:
        return

    terminal = result if isinstance(result, dict) else {}
    interrupted = terminal.get("interrupted") is True
    completed = terminal.get("completed") is True
    failed = terminal.get("failed") is True
    reason = str(
        terminal.get("turn_exit_reason") or terminal.get("failure_reason") or ""
    )
    if error is not None:
        interrupted = isinstance(error, (KeyboardInterrupt, InterruptedError)) or (
            type(error).__name__ == "CancelledError"
        )
        timed_out = isinstance(error, TimeoutError)
        completed = False
        failed = not interrupted
        if interrupted:
            reason = "interrupted_by_user"
        elif timed_out:
            reason = "timed_out"
        else:
            reason = "system_aborted"
    elif not reason:
        reason = "failed" if failed else "unknown"

    runtime._safe(
        runtime.finish_task,
        {
            "session_id": session_id,
            "task_id": task_id,
            "platform": platform,
            "completed": completed,
            "failed": failed,
            "interrupted": interrupted,
            "turn_exit_reason": reason,
        },
    )


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()
            _RUNTIMES.pop(profile_key, None)
        if runtime is _RUNTIME_FAILED and not retry_failed:
            return None
        if runtime is _RUNTIME_FAILED:
            _RUNTIMES.pop(profile_key, None)
        try:
            runtime = _Runtime(host=host)
        except Exception:
            logger.warning("Hermes shared metrics initialization failed", exc_info=True)
            _RUNTIMES[profile_key] = _RUNTIME_FAILED
            return None
        _RUNTIMES[profile_key] = runtime
        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()
