Files
hermes-agent/tests/tui_gateway/test_project_recall_rpc.py

222 lines
11 KiB
Python

"""Project recall RPCs authorize live callers before reading real SQLite history."""
from pathlib import Path
from queue import Queue
from types import SimpleNamespace
import pytest
@pytest.fixture
def runtime(tmp_path, monkeypatch):
from tui_gateway import server
from hermes_state_registry import acquire
home = tmp_path / ".hermes"
home.mkdir()
monkeypatch.setattr(Path, "home", lambda: tmp_path)
monkeypatch.setenv("HERMES_HOME", str(home))
monkeypatch.setattr(server, "_hermes_home", home)
monkeypatch.setattr(server, "_served_profile_homes", set())
root = tmp_path / "project"
root.mkdir()
other = tmp_path / "foreign-project"
other.mkdir()
db = acquire(home / "state.db")
for sid, cwd in (("current", root), ("own", root), ("foreign", other), ("home", None)):
db.create_session(sid, source="desktop", cwd=str(cwd) if cwd else None)
mid = db.append_message("own", role="user", content="recallneedle original evidence")
db.append_message("foreign", role="user", content="recallneedle FOREIGN SECRET")
responses = Queue()
transport = SimpleNamespace(write=lambda frame: responses.put(frame))
owner = {"session_key": "current", "profile_home": None, "transport": transport}
monkeypatch.setattr(server, "_sessions", {"live-ui": owner})
def call(method, *, via=transport, **params):
response = server.dispatch({"jsonrpc": "2.0", "id": 1, "method": method,
"params": {"session_id": "live-ui", **params}}, transport=via)
return responses.get(timeout=15) if response is None else response
try:
yield SimpleNamespace(server=server, db=db, call=call, owner=owner,
transport=transport, home=home, root=root, mid=mid)
finally:
db.close()
def test_search_uses_authorized_stored_identity_and_project_sqlite(runtime):
result = runtime.call("project.recall.search", query="recallneedle")["result"]
assert result["success"]
assert [row["session_id"] for row in result["results"]] == ["own"]
assert result["results"][0]["match_message_id"] == runtime.mid
assert result["results"][0]["messages"][0]["content"] == "recallneedle original evidence"
assert result["scope"]["status"] == "ready"
assert result["scope"]["label"] == runtime.root.name
assert "allowed_session_ids" not in result["scope"]
assert "scanned_sessions" not in result["scope"]["coverage"]
def test_read_preserves_source_and_status_only_counts_this_project(runtime):
result = runtime.call("project.recall.read", source_session_id="own",
around_message_id=runtime.mid, content_offset=3, content_length=7)["result"]
assert result["session_id"] == "own"
assert result["messages"][0]["id"] == runtime.mid
assert result["messages"][0]["content"] == "allneed"
assert result["messages"][0]["content_offset"] == 3
status = runtime.call("project.recall.status")["result"]
assert status["success"] and status["scope"]["status"] == "ready"
assert status["session_count"] == 2
assert "scanned_sessions" not in status["scope"]["coverage"]
assert "results" not in status and "messages" not in status
@pytest.mark.parametrize("method", ["search", "read", "status"])
def test_transport_authority_and_no_project_fail_closed(runtime, method):
rpc = "project.recall." + method
params = {"query": "recallneedle"} if method == "search" else (
{"source_session_id": "own"} if method == "read" else {})
foreign = SimpleNamespace(write=lambda frame: True)
runtime.server._sessions["foreign-live"] = {**runtime.owner, "transport": foreign}
assert runtime.call(rpc, session_id="foreign-live", **params)["error"]["code"] == 4001
assert runtime.call(rpc, session_id="current", **params)["error"]["code"] == 4001
runtime.owner["session_key"] = "home"
result = runtime.call(rpc, **params)["result"]
assert result["success"] is False and result["status"] == "scope_unresolved"
assert result["scope"]["status"] == "scope_unresolved"
assert not result.get("results") and not result.get("messages")
@pytest.mark.parametrize("profile", ["../escape", "/tmp", "", "Default", {}, "missing"])
def test_invalid_profile_is_rejected_without_fallback(runtime, profile):
result = runtime.call("project.recall.search", profile=profile, query="recallneedle")
assert result["error"]["code"] == 4000
def test_target_cannot_escape_project_or_override_trusted_identity(runtime):
result = runtime.call("project.recall.read", source_session_id="foreign")["result"]
assert result["status"] == "access_denied" and not result.get("messages")
result = runtime.call("project.recall.read", source_session_id="other/own")["result"]
assert result["status"] == "access_denied"
result = runtime.call("project.recall.search", current_session_id="foreign", query="recallneedle")
assert result["error"]["code"] == 4000
def test_profile_context_comes_from_owner_and_never_falls_back(runtime, monkeypatch):
from hermes_state_registry import acquire
import hermes_state_registry
other_home = runtime.home / "profiles" / "research"
other_home.mkdir(parents=True)
db = acquire(other_home / "state.db")
try:
for sid in ("current", "profile-own"):
db.create_session(sid, source="desktop", cwd=str(runtime.root))
db.append_message("profile-own", role="user", content="recallneedle SECOND PROFILE")
runtime.owner["profile_home"] = str(other_home)
result = runtime.call("project.recall.search", query="recallneedle")["result"]
assert [row["session_id"] for row in result["results"]] == ["profile-own"]
assert runtime.call("project.recall.search", profile="default")["error"]["code"] == 4001
assert runtime.call("project.recall.status", profile="research")["result"]["session_count"] == 2
def unavailable(path):
assert Path(path).resolve() == (other_home / "state.db").resolve()
raise OSError("unavailable")
monkeypatch.setattr(hermes_state_registry, "acquire", unavailable)
assert runtime.call("project.recall.search", query="recallneedle")["error"]["code"] == 5031
finally:
db.close()
@pytest.mark.parametrize("params", [{"window": 0}, {"limit": True}, {"query": []},
{"search_cursor": 10}, {"scope": "legacy_profile"}])
def test_search_validates_rpc_parameters(runtime, params):
assert runtime.call("project.recall.search", **params)["error"]["code"] == 4000
def test_live_generation_revocation_during_query_drops_result(runtime, monkeypatch):
from tools import session_search_project
original = session_search_project.project_session_search
def revoke(*args, **kwargs):
result = original(*args, **kwargs)
runtime.server._sessions["live-ui"] = {**runtime.owner}
return result
monkeypatch.setattr(session_search_project, "project_session_search", revoke)
assert runtime.call("project.recall.search", query="recallneedle")["error"]["code"] == 4001
def test_search_cursor_and_read_scan_frontiers_round_trip(runtime):
runtime.db.create_session("own-next", source="desktop", cwd=str(runtime.root))
runtime.db.append_message("own-next", role="user", content="recallneedle next source")
first = runtime.call("project.recall.search", query="recallneedle", limit=1, sort="oldest")["result"]
second = runtime.call("project.recall.search", query="recallneedle", limit=1, sort="oldest",
search_cursor=first["next_search_cursor"])["result"]
assert {row["session_id"] for page in (first, second) for row in page["results"]} == {"own", "own-next"}
ids = [runtime.mid]
for index in range(35):
ids.append(runtime.db.append_message("own", role="assistant", content=f"source line {index}"))
page = runtime.call("project.recall.read", source_session_id="own")["result"]
later = runtime.call("project.recall.read", source_session_id="own",
after_message_id=page["next_after_message_id"])["result"]
assert [m["id"] for p in (page, later) for m in p["messages"]] == ids
around = runtime.call("project.recall.read", source_session_id="own", around_message_id=ids[20],
before_message_id=ids[15], after_message_id=ids[25], window=2)["result"]
assert [m["id"] for m in around["messages"]] == ids[13:15] + [ids[20]] + ids[26:28]
def test_compressed_live_agent_id_takes_precedence_over_stale_session_key(runtime):
runtime.owner["session_key"] = "foreign"
runtime.owner["agent"] = SimpleNamespace(session_id="current")
result = runtime.call("project.recall.search", query="recallneedle")["result"]
assert [row["session_id"] for row in result["results"]] == ["own"]
def test_scope_move_during_query_drops_all_history(runtime, monkeypatch):
from tools import session_search_project
original = session_search_project.project_session_search
def move(*args, **kwargs):
result = original(*args, **kwargs)
runtime.db._conn.execute("UPDATE sessions SET cwd=NULL WHERE id='current'")
runtime.db._conn.commit()
return result
monkeypatch.setattr(session_search_project, "project_session_search", move)
result = runtime.call("project.recall.search", query="recallneedle")["result"]
assert result["status"] == "scope_changed" and not result.get("results")
@pytest.mark.parametrize('method,params', [
('search', {'search_cursor': 'cursor-without-query'}),
('search', {'limit': 11}),
('read', {'source_session_id': 'own', 'content_offset': 4}),
('read', {'source_session_id': 'own', 'content_length': 4001}),
])
def test_rpc_rejects_incomplete_combinations_and_unified_budget(runtime, method, params):
reply = runtime.call('project.recall.' + method, **params)
assert reply['error']['code'] == 4000
def test_rpc_revalidates_returned_sources_after_inner_search(runtime, monkeypatch):
from tools import session_search_project
original = session_search_project.project_session_search
def revoke(*args, **kwargs):
result = original(*args, **kwargs)
runtime.db._conn.execute('DELETE FROM messages WHERE id=?', (runtime.mid,))
runtime.db._conn.commit()
return result
monkeypatch.setattr(session_search_project, 'project_session_search', revoke)
result = runtime.call('project.recall.search', query='recallneedle')['result']
assert result['status'] == 'source_revoked'
assert not result.get('results')
def test_unbound_transport_cannot_use_stored_session_identity(runtime):
result = runtime.server.handle_request({"id": 1, "method": "project.recall.status",
"params": {"session_id": "live-ui"}})
assert result["error"]["code"] == 4001