99 lines
3.7 KiB
Python
99 lines
3.7 KiB
Python
"""Nested code RPC must retain the caller, not the requested read target."""
|
|
|
|
import json
|
|
from concurrent.futures import ThreadPoolExecutor
|
|
from contextlib import contextmanager
|
|
from dataclasses import replace
|
|
|
|
import pytest
|
|
|
|
from tools.approval_context import (
|
|
reset_current_observability_context,
|
|
set_current_observability_context,
|
|
)
|
|
from tools.code_execution_rpc import _default_dispatch
|
|
from tools.thread_context import propagate_context_to_thread
|
|
|
|
|
|
@contextmanager
|
|
def caller(session_id):
|
|
tokens = set_current_observability_context(session_id=session_id)
|
|
try:
|
|
yield
|
|
finally:
|
|
reset_current_observability_context(tokens)
|
|
|
|
|
|
@pytest.fixture
|
|
def identity_handler(monkeypatch):
|
|
from model_tools import registry
|
|
|
|
monkeypatch.setenv("HERMES_SESSION_ID", "foreign")
|
|
# Keep the actual dispatcher and registry; replace only the network leaf.
|
|
def identity(args, **kwargs):
|
|
return json.dumps({"args": args, "session_id": kwargs.get("session_id"),
|
|
"task_id": kwargs.get("task_id")})
|
|
|
|
monkeypatch.setitem(registry._tools, "web_search", replace(
|
|
registry._tools["web_search"], handler=identity))
|
|
|
|
|
|
def test_rpc_dispatch_preserves_trusted_caller_across_worker_context(identity_handler):
|
|
def read_history():
|
|
dispatch = _default_dispatch("not-a-session")
|
|
return json.loads(dispatch("web_search", {
|
|
"query": "history", "session_id": "foreign", "current_session_id": "foreign",
|
|
}))
|
|
|
|
with caller("caller"):
|
|
trusted_worker = propagate_context_to_thread(read_history)
|
|
with caller(""):
|
|
anonymous_worker = propagate_context_to_thread(read_history)
|
|
with caller("foreign"), ThreadPoolExecutor(max_workers=2) as pool:
|
|
trusted = pool.submit(trusted_worker)
|
|
anonymous = pool.submit(anonymous_worker)
|
|
own = trusted.result(timeout=20)
|
|
denied = anonymous.result(timeout=20)
|
|
|
|
assert own["session_id"] == "caller", own
|
|
assert own["task_id"] == "not-a-session"
|
|
assert own["args"]["session_id"] == "foreign"
|
|
assert not denied["session_id"], denied
|
|
|
|
|
|
def test_reused_kernel_alias_uses_each_outer_calls_identity(identity_handler, monkeypatch):
|
|
from model_tools import handle_function_call
|
|
from tools import code_execution_tool
|
|
from tools.code_kernel import _KERNELS, shutdown_all_kernels
|
|
|
|
monkeypatch.setenv("TERMINAL_ENV", "local")
|
|
monkeypatch.setattr(code_execution_tool, "_load_config", lambda: {
|
|
"mode": "strict", "timeout": 30,
|
|
})
|
|
code = "import json; print(json.dumps(alias(query='history')))"
|
|
results = []
|
|
try:
|
|
for session_id, source in [
|
|
("first-caller", "from hermes_tools import web_search as alias\n" + code),
|
|
("second-caller", code),
|
|
(None, code),
|
|
]:
|
|
result = json.loads(handle_function_call(
|
|
"execute_code", {"code": source}, task_id="kernel-identity-task",
|
|
session_id=session_id, enabled_tools=["web_search"],
|
|
))
|
|
assert result["status"] == "success", result
|
|
results.append(result)
|
|
for kernel in _KERNELS.values():
|
|
assert kernel.cell_authority is not None
|
|
retired = json.loads(kernel.cell_authority.dispatch("web_search", {}))
|
|
assert "No active execute_code cell" in retired["error"]
|
|
finally:
|
|
shutdown_all_kernels()
|
|
|
|
assert results[1]["kernel"]["reused"]
|
|
assert results[2]["kernel"]["reused"]
|
|
for result, expected in zip(results, ["first-caller", "second-caller", None]):
|
|
nested = json.loads(result["output"])
|
|
assert nested["session_id"] == expected, nested
|
|
assert nested["task_id"] == "kernel-identity-task" |