refactor(hermes_cli): relay_shared_metrics — fold model-call error/end into one updater, _task_pair lookup

This commit is contained in:
Teknium
2026-09-02 21:15:27 -07:00
parent c9c13fa8f3
commit 50ed27f263
@@ -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: