diff --git a/hermes_cli/observability/relay_shared_metrics.py b/hermes_cli/observability/relay_shared_metrics.py index 7c3798e3aa..bed781e139 100644 --- a/hermes_cli/observability/relay_shared_metrics.py +++ b/hermes_cli/observability/relay_shared_metrics.py @@ -272,10 +272,10 @@ class _Runtime: 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: + request_id = _text(event, "api_request_id") + if not request_id: return - _, request_id = model_call_key + model_call_key = (task_id, request_id) fields = model_call_fields(event) retry_ordinal = _retry_ordinal(event) with session.lock: @@ -317,15 +317,26 @@ class _Runtime: def record_model_call_error(self, event: dict[str, Any]) -> None: """Retain the latest attempt error without closing the logical call.""" + self._update_model_call(event, finish=False) + + def end_model_call(self, event: dict[str, Any]) -> None: + self._update_model_call(event, finish=True) + + def _update_model_call(self, event: dict[str, Any], *, finish: bool) -> None: + """Refresh the located model call's fields from ``event``, optionally closing it.""" session = self._any_session(event) if session is None: return with session.lock: if session.closing: return - located = self._model_call_for(session, event) - if located is not None: - located[1].fields = model_call_fields(event) + 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 = model_call_fields(event) + if finish: + self._finish_model_call(session, model_call_key) def start_tool_call(self, event: dict[str, Any]) -> None: """Open one privacy-safe Relay tool lifecycle under its task.""" @@ -398,9 +409,8 @@ class _Runtime: # that reused the provider-local ID. return if matching_keys: - key = matching_keys[0] - identity = key[1:] - tool_call = session.tool_calls.pop(key) + 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: @@ -419,9 +429,8 @@ class _Runtime: return session_id, task_id = _text(event, "session_id"), _text(event, "task_id") - session = self._task_session(event, allow_task_id_fallback=not session_id) + session, task = self._task_pair(event, allow_task_id_fallback=not session_id) if session is not None: - task = session.tasks.get(task_id) if task is None: return with session.lock: @@ -439,20 +448,6 @@ class _Runtime: 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._any_session(event) - if session is None: - return - with session.lock: - if session.closing: - return - located = self._model_call_for(session, event) - if located is None: - return - model_call_key, model_call = located - model_call.fields = model_call_fields(event) - self._finish_model_call(session, model_call_key) - def finish_task(self, event: dict[str, Any]) -> None: """Close one task scope exactly once with bounded terminal fields.""" session = self._any_session(event) @@ -560,21 +555,18 @@ class _Runtime: self, event: dict[str, Any], *, start: bool ) -> tuple[_MetricsSession | None, _TaskRun | None]: """Resolve (session, task) for a task-scoped hook, optionally opening the task.""" - session = self._task_session(event, allow_task_id_fallback=True) - task = session.tasks.get(_text(event, "task_id")) if session is not None else None + session, task = self._task_pair(event, allow_task_id_fallback=True) 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 _model_call_for( - self, session: _MetricsSession, event: dict[str, Any] - ) -> tuple[tuple[str, str], _ModelCall] | None: - model_call_key = self._existing_model_call_key(session, event) - if model_call_key is None: - return None - model_call = session.model_calls.get(model_call_key) - return None if model_call is None else (model_call_key, model_call) + def _task_pair( + self, event: dict[str, Any], **lookup: Any + ) -> tuple[_MetricsSession | None, _TaskRun | None]: + session = self._task_session(event, **lookup) + task = session.tasks.get(_text(event, "task_id")) if session is not None else None + return session, task def _run_scoped( self, session: _MetricsSession, task: _TaskRun | None, callback: Callable[..., Any], @@ -678,14 +670,13 @@ class _Runtime: """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 + 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 = self._task_session(event) - task = session.tasks.get(_text(event, "task_id")) if session is not None else None + session, task = self._task_pair(event) if task is not None: return session, task @@ -763,22 +754,19 @@ class _Runtime: self._finish_model_call(session, model_call_key) @staticmethod - def _new_model_call_key(event: dict[str, Any]) -> tuple[str, str] | None: - request_id = _text(event, "api_request_id") - return (_text(event, "task_id"), request_id) if request_id else None - - @classmethod def _existing_model_call_key( - cls, session: _MetricsSession, event: dict[str, Any] + session: _MetricsSession, event: dict[str, Any] ) -> tuple[str, str] | None: - key = cls._new_model_call_key(event) - if key is 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: return key if key[0]: return None - candidates = [candidate for candidate in session.model_calls if candidate[1] == key[1]] + 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: