fix(observability): harden tool lifecycle metrics
Signed-off-by: Alex Fournier <afournier@nvidia.com>
This commit is contained in:
parent
4ad78a98fb
commit
8502e464a8
|
|
@ -381,12 +381,20 @@ class _Runtime:
|
|||
identity = self._tool_call_identity(event)
|
||||
tool_call = session.tool_calls.get((task.task_id, *identity))
|
||||
if tool_call is None:
|
||||
matches = [
|
||||
candidate
|
||||
for key, candidate in session.tool_calls.items()
|
||||
if key[0] == task.task_id and key[-1] == tool_call_id
|
||||
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 = matches[0] if len(matches) == 1 else None
|
||||
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"
|
||||
|
|
@ -414,15 +422,42 @@ class _Runtime:
|
|||
return
|
||||
self._remember_turn(session, task, event)
|
||||
if tool_call_id:
|
||||
identity = self._tool_call_identity(event)
|
||||
if identity in task.completed_tool_call_ids:
|
||||
observed_identity = self._tool_call_identity(event)
|
||||
if observed_identity in task.completed_tool_call_ids:
|
||||
return
|
||||
task.completed_tool_call_ids.add(identity)
|
||||
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)
|
||||
tool_call = session.tool_calls.pop(
|
||||
(task_id, *identity),
|
||||
None,
|
||||
)
|
||||
else:
|
||||
task.unidentified_tool_calls += 1
|
||||
tool_call = None
|
||||
|
|
@ -659,6 +694,25 @@ class _Runtime:
|
|||
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,
|
||||
|
|
@ -999,9 +1053,9 @@ def _with_runtime_toolset(event: dict[str, Any]) -> dict[str, Any]:
|
|||
if not tool_name:
|
||||
return event
|
||||
try:
|
||||
from tools.registry import registry
|
||||
from model_tools import get_toolset_for_tool
|
||||
|
||||
toolset = registry.get_toolset_for_tool(tool_name)
|
||||
toolset = get_toolset_for_tool(tool_name)
|
||||
except Exception:
|
||||
toolset = None
|
||||
return {**event, "toolset": toolset or "other"}
|
||||
|
|
|
|||
|
|
@ -2614,6 +2614,137 @@ def test_reused_tool_call_id_is_counted_for_each_provider_request(direct_runtime
|
|||
assert task_end[2]["output"]["tool_call_count_bucket"] == "2"
|
||||
|
||||
|
||||
def test_partial_terminal_context_reuses_the_pending_tool_span(direct_runtime):
|
||||
base = {
|
||||
"session_id": "s1",
|
||||
"task_id": "t1",
|
||||
"turn_id": "turn-1",
|
||||
"api_request_id": "request-1",
|
||||
"platform": "cli",
|
||||
"tool_call_id": "tool-1",
|
||||
"tool_name": "terminal",
|
||||
}
|
||||
lifecycle.invoke_hook("pre_llm_call", **base)
|
||||
lifecycle.invoke_hook("pre_tool_call", **base)
|
||||
lifecycle.invoke_hook(
|
||||
"post_tool_call",
|
||||
**{key: value for key, value in base.items() if key != "api_request_id"},
|
||||
result={"output": "private"},
|
||||
status="ok",
|
||||
)
|
||||
lifecycle.invoke_hook(
|
||||
"on_session_end",
|
||||
**base,
|
||||
completed=True,
|
||||
failed=False,
|
||||
interrupted=False,
|
||||
turn_exit_reason="text_response(stop)",
|
||||
)
|
||||
lifecycle.finalize_session(session_id="s1")
|
||||
|
||||
[tool_end] = [
|
||||
event for event in direct_runtime.events if event[0] == "tool.call_end"
|
||||
]
|
||||
assert tool_end[2]["outcome"] == "success"
|
||||
[task_end] = [
|
||||
event
|
||||
for event in direct_runtime.events
|
||||
if event[0] == "scope.pop" and event[1][1] == "hermes.task_run"
|
||||
]
|
||||
assert task_end[2]["output"]["tool_call_count_bucket"] == "1"
|
||||
|
||||
|
||||
def test_partial_terminal_variants_do_not_double_count_a_completed_call(
|
||||
direct_runtime,
|
||||
):
|
||||
base = {
|
||||
"session_id": "s1",
|
||||
"task_id": "t1",
|
||||
"turn_id": "turn-1",
|
||||
"api_request_id": "request-1",
|
||||
"platform": "cli",
|
||||
"tool_call_id": "tool-1",
|
||||
"tool_name": "terminal",
|
||||
}
|
||||
lifecycle.invoke_hook("pre_llm_call", **base)
|
||||
lifecycle.invoke_hook("pre_tool_call", **base)
|
||||
for omitted_field in ("api_request_id", "turn_id"):
|
||||
lifecycle.invoke_hook(
|
||||
"post_tool_call",
|
||||
**{key: value for key, value in base.items() if key != omitted_field},
|
||||
result={"output": "private"},
|
||||
status="ok",
|
||||
)
|
||||
lifecycle.invoke_hook(
|
||||
"on_session_end",
|
||||
**base,
|
||||
completed=True,
|
||||
failed=False,
|
||||
interrupted=False,
|
||||
turn_exit_reason="text_response(stop)",
|
||||
)
|
||||
lifecycle.finalize_session(session_id="s1")
|
||||
|
||||
tool_ends = [
|
||||
event for event in direct_runtime.events if event[0] == "tool.call_end"
|
||||
]
|
||||
assert len(tool_ends) == 1
|
||||
[task_end] = [
|
||||
event
|
||||
for event in direct_runtime.events
|
||||
if event[0] == "scope.pop" and event[1][1] == "hermes.task_run"
|
||||
]
|
||||
assert task_end[2]["output"]["tool_call_count_bucket"] == "1"
|
||||
|
||||
|
||||
def test_ambiguous_partial_terminal_does_not_create_a_phantom_tool_span(
|
||||
direct_runtime,
|
||||
):
|
||||
base = {
|
||||
"session_id": "s1",
|
||||
"task_id": "t1",
|
||||
"turn_id": "turn-1",
|
||||
"platform": "cli",
|
||||
"tool_call_id": "provider-reused-id",
|
||||
"tool_name": "terminal",
|
||||
}
|
||||
lifecycle.invoke_hook("pre_llm_call", **base)
|
||||
for api_request_id in ("request-1", "request-2"):
|
||||
lifecycle.invoke_hook(
|
||||
"pre_tool_call",
|
||||
**base,
|
||||
api_request_id=api_request_id,
|
||||
)
|
||||
|
||||
lifecycle.invoke_hook(
|
||||
"post_tool_call",
|
||||
**base,
|
||||
result={"output": "ambiguous-private-result"},
|
||||
status="ok",
|
||||
)
|
||||
lifecycle.invoke_hook(
|
||||
"on_session_end",
|
||||
**base,
|
||||
completed=False,
|
||||
failed=True,
|
||||
interrupted=False,
|
||||
turn_exit_reason="system_aborted",
|
||||
)
|
||||
lifecycle.finalize_session(session_id="s1")
|
||||
|
||||
tool_ends = [
|
||||
event for event in direct_runtime.events if event[0] == "tool.call_end"
|
||||
]
|
||||
assert len(tool_ends) == 2
|
||||
assert all(event[2]["outcome"] == "failed" for event in tool_ends)
|
||||
[task_end] = [
|
||||
event
|
||||
for event in direct_runtime.events
|
||||
if event[0] == "scope.pop" and event[1][1] == "hermes.task_run"
|
||||
]
|
||||
assert task_end[2]["output"]["tool_call_count_bucket"] == "2"
|
||||
|
||||
|
||||
def test_reused_task_id_starts_a_new_run_for_each_turn(direct_runtime):
|
||||
for turn_id in ("turn-1", "turn-2"):
|
||||
base = {
|
||||
|
|
@ -2920,10 +3051,10 @@ def test_tool_category_comes_from_runtime_registry_metadata(
|
|||
direct_runtime,
|
||||
monkeypatch,
|
||||
):
|
||||
from tools.registry import registry
|
||||
import model_tools
|
||||
|
||||
monkeypatch.setattr(
|
||||
registry,
|
||||
model_tools,
|
||||
"get_toolset_for_tool",
|
||||
lambda name: "terminal" if name == "runtime_only_tool" else None,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -2427,6 +2427,43 @@ class TestApprovalTimeoutIsNotConsent:
|
|||
f"hook choice should be 'timeout' on no-response, got {last_post.get('choice')!r}"
|
||||
)
|
||||
|
||||
def test_notify_failure_emits_post_hook_and_cleans_up(self, monkeypatch):
|
||||
"""A failed notification still terminates the approval lifecycle."""
|
||||
from tools import approval as mod
|
||||
|
||||
hook_calls = []
|
||||
|
||||
def _capture(event_name, **kwargs):
|
||||
hook_calls.append((event_name, kwargs))
|
||||
|
||||
monkeypatch.setattr(mod, "_fire_approval_hook", _capture)
|
||||
|
||||
def _fail_notify(_data):
|
||||
raise RuntimeError("private gateway failure")
|
||||
|
||||
decision = mod._await_gateway_decision(
|
||||
self.SESSION_KEY,
|
||||
_fail_notify,
|
||||
{
|
||||
"command": "redacted-command",
|
||||
"description": "redacted-description",
|
||||
"pattern_key": "dangerous",
|
||||
"pattern_keys": ["dangerous"],
|
||||
},
|
||||
)
|
||||
|
||||
assert decision == {
|
||||
"resolved": False,
|
||||
"choice": None,
|
||||
"notify_failed": True,
|
||||
}
|
||||
assert self.SESSION_KEY not in mod._gateway_queues
|
||||
assert [name for name, _ in hook_calls] == [
|
||||
"pre_approval_request",
|
||||
"post_approval_response",
|
||||
]
|
||||
assert hook_calls[-1][1]["choice"] == "notify_failed"
|
||||
|
||||
|
||||
class TestTirithImportErrorFailOpenPolicy:
|
||||
"""Regression guard for #20733.
|
||||
|
|
|
|||
|
|
@ -3286,6 +3286,16 @@ def _await_gateway_decision(session_key: str, notify_cb, approval_data: dict,
|
|||
except Exception as exc:
|
||||
logger.warning("Gateway approval notify failed: %s", exc)
|
||||
_drop_entry()
|
||||
_fire_approval_hook(
|
||||
"post_approval_response",
|
||||
command=command,
|
||||
description=description,
|
||||
pattern_key=primary_key,
|
||||
pattern_keys=list(all_keys),
|
||||
session_key=session_key,
|
||||
surface=surface,
|
||||
choice="notify_failed",
|
||||
)
|
||||
return {"resolved": False, "choice": None, "notify_failed": True}
|
||||
|
||||
# Block until the user responds or the canonical approval timeout elapses
|
||||
|
|
|
|||
Loading…
Reference in New Issue