mirror of
https://github.com/turnstonelabs/turnstone.git
synced 2026-08-12 23:12:23 -06:00
71c34839d9
Per-user MCP credential resolution was bound once at session construction to the persisted workstream owner, so on a shared workstream every sender executed oauth_user tools under the creator's tokens (and saw the creator's tool catalog). Bind the authenticated initiator of each turn (send + retry paths) as the session's acting user: dispatch, catalog merge, visibility gates, and consent flows now follow whoever is driving, with the owner as fallback for CLI / eval / scheduled / internal turns. Rebinding swaps the user-scoped tool/resource/prompt listeners (identity is the (user_id, callback) pair), fire-and-forget primes the acting user's pools, and rebuilds the merged tool list. Prepared tool items pin the identity at prepare time so an item pending approval executes under the user whose turn requested it, not whoever binds later. Queued mid-turn interjections deliberately do not rebind (no mid-turn credential switch).
1338 lines
52 KiB
Python
1338 lines
52 KiB
Python
"""Tests for workstream persistence and resume functionality."""
|
|
|
|
from unittest.mock import MagicMock
|
|
|
|
import sqlalchemy as sa
|
|
|
|
from turnstone.core.memory import (
|
|
delete_workstream,
|
|
list_workstreams_with_history,
|
|
load_messages,
|
|
load_workstream_config,
|
|
prune_workstreams,
|
|
register_workstream,
|
|
resolve_workstream,
|
|
save_message,
|
|
save_workstream_config,
|
|
set_workstream_alias,
|
|
update_workstream_title,
|
|
)
|
|
from turnstone.core.session import ChatSession
|
|
from turnstone.core.storage import get_storage
|
|
from turnstone.core.trajectory import turn_to_dict
|
|
|
|
# ── Workstream registration ───────────────────────────────────────────
|
|
|
|
|
|
class TestRegisterWorkstream:
|
|
def test_register_creates_row(self, tmp_db):
|
|
register_workstream("abc123")
|
|
# Workstream exists in DB (resolve works) even without messages
|
|
assert resolve_workstream("abc123") == "abc123"
|
|
|
|
def test_register_with_title(self, tmp_db):
|
|
register_workstream("abc123", name="My Workstream")
|
|
save_message("abc123", "user", "hello")
|
|
rows = list_workstreams_with_history()
|
|
assert rows[0][2] is None # title column (name is separate)
|
|
|
|
def test_register_idempotent(self, tmp_db):
|
|
register_workstream("abc123")
|
|
update_workstream_title("abc123", "First")
|
|
register_workstream("abc123") # should be ignored
|
|
update_workstream_title("abc123", "First") # title is set via update
|
|
save_message("abc123", "user", "hello")
|
|
rows = list_workstreams_with_history()
|
|
assert len(rows) == 1
|
|
assert rows[0][2] == "First" # title preserved
|
|
|
|
def test_update_title(self, tmp_db):
|
|
register_workstream("abc123")
|
|
update_workstream_title("abc123", "New Title")
|
|
save_message("abc123", "user", "hello")
|
|
rows = list_workstreams_with_history()
|
|
assert rows[0][2] == "New Title"
|
|
|
|
|
|
# ── Workstream alias ──────────────────────────────────────────────────
|
|
|
|
|
|
class TestWorkstreamAlias:
|
|
def test_set_alias(self, tmp_db):
|
|
register_workstream("abc123")
|
|
assert set_workstream_alias("abc123", "my-session") is True
|
|
save_message("abc123", "user", "hello")
|
|
rows = list_workstreams_with_history()
|
|
assert rows[0][1] == "my-session" # alias
|
|
|
|
def test_alias_conflict(self, tmp_db):
|
|
register_workstream("abc123")
|
|
register_workstream("def456")
|
|
set_workstream_alias("abc123", "taken")
|
|
assert set_workstream_alias("def456", "taken") is False
|
|
|
|
def test_alias_same_workstream_ok(self, tmp_db):
|
|
register_workstream("abc123")
|
|
set_workstream_alias("abc123", "mine")
|
|
assert set_workstream_alias("abc123", "mine") is True # no-op, same workstream
|
|
|
|
|
|
# ── Workstream resolution ─────────────────────────────────────────────
|
|
|
|
|
|
class TestResolveWorkstream:
|
|
def test_resolve_by_alias(self, tmp_db):
|
|
register_workstream("abc123")
|
|
set_workstream_alias("abc123", "my-alias")
|
|
assert resolve_workstream("my-alias") == "abc123"
|
|
|
|
def test_resolve_by_exact_id(self, tmp_db):
|
|
register_workstream("abc123def456")
|
|
assert resolve_workstream("abc123def456") == "abc123def456"
|
|
|
|
def test_resolve_by_prefix(self, tmp_db):
|
|
register_workstream("abc123def456")
|
|
assert resolve_workstream("abc123") == "abc123def456"
|
|
|
|
def test_resolve_prefix_ambiguous(self, tmp_db):
|
|
register_workstream("abc123aaaaaa")
|
|
register_workstream("abc123bbbbbb")
|
|
# Ambiguous prefix should return None
|
|
assert resolve_workstream("abc123") is None
|
|
|
|
def test_resolve_not_found(self, tmp_db):
|
|
assert resolve_workstream("nonexistent") is None
|
|
|
|
|
|
# ── List workstreams with history ──────────────────────────────────────
|
|
|
|
|
|
class TestListWorkstreamsWithHistory:
|
|
def test_empty(self, tmp_db):
|
|
assert list_workstreams_with_history() == []
|
|
|
|
def test_ordered_by_updated(self, tmp_db):
|
|
register_workstream("first")
|
|
save_message("first", "user", "hello")
|
|
# Force an older timestamp so ordering is deterministic
|
|
engine = get_storage()._engine # noqa: SLF001
|
|
with engine.connect() as conn:
|
|
conn.execute(
|
|
sa.text("UPDATE workstreams SET updated = '2020-01-01' WHERE ws_id = 'first'")
|
|
)
|
|
conn.commit()
|
|
register_workstream("second")
|
|
save_message("second", "user", "hello")
|
|
# second is more recent
|
|
rows = list_workstreams_with_history()
|
|
assert rows[0][0] == "second"
|
|
assert rows[1][0] == "first"
|
|
|
|
def test_includes_message_count(self, tmp_db):
|
|
register_workstream("sess1")
|
|
save_message("sess1", "user", "hello")
|
|
save_message("sess1", "assistant", "hi")
|
|
rows = list_workstreams_with_history()
|
|
assert rows[0][6] == 2 # msg_count (after ws_id, alias, title, name, created, updated)
|
|
|
|
def test_respects_limit(self, tmp_db):
|
|
for i in range(5):
|
|
register_workstream(f"sess{i}")
|
|
save_message(f"sess{i}", "user", "hello")
|
|
rows = list_workstreams_with_history(limit=3)
|
|
assert len(rows) == 3
|
|
|
|
|
|
# ── Load messages ─────────────────────────────────────────────────────
|
|
|
|
|
|
class TestLoadMessages:
|
|
def test_simple_user_assistant(self, tmp_db):
|
|
save_message("s1", "user", "hello")
|
|
save_message("s1", "assistant", "hi there")
|
|
msgs = load_messages("s1")
|
|
assert len(msgs) == 2
|
|
assert msgs[0] == {"role": "user", "content": "hello"}
|
|
assert msgs[1] == {"role": "assistant", "content": "hi there"}
|
|
|
|
def test_tool_calls_with_ids(self, tmp_db):
|
|
import json
|
|
|
|
tc_json = json.dumps(
|
|
[
|
|
{
|
|
"id": "call_abc",
|
|
"type": "function",
|
|
"function": {"name": "bash", "arguments": '{"command":"ls"}'},
|
|
}
|
|
]
|
|
)
|
|
save_message("s1", "user", "run ls")
|
|
save_message("s1", "assistant", "Let me check.", tool_calls=tc_json)
|
|
save_message("s1", "tool", "file1.txt\nfile2.txt", "bash", tool_call_id="call_abc")
|
|
msgs = load_messages("s1")
|
|
assert len(msgs) == 3 # user, assistant+tool_calls, tool
|
|
# Assistant should have content and tool_calls
|
|
assert msgs[1]["role"] == "assistant"
|
|
assert msgs[1]["content"] == "Let me check."
|
|
assert len(msgs[1]["tool_calls"]) == 1
|
|
assert msgs[1]["tool_calls"][0]["id"] == "call_abc"
|
|
assert msgs[1]["tool_calls"][0]["function"]["name"] == "bash"
|
|
# Tool result
|
|
assert msgs[2]["role"] == "tool"
|
|
assert msgs[2]["tool_call_id"] == "call_abc"
|
|
assert msgs[2]["content"] == "file1.txt\nfile2.txt"
|
|
|
|
def test_parallel_tool_calls(self, tmp_db):
|
|
import json
|
|
|
|
tc_json = json.dumps(
|
|
[
|
|
{
|
|
"id": "call_1",
|
|
"type": "function",
|
|
"function": {"name": "search", "arguments": '{"query":"a"}'},
|
|
},
|
|
{
|
|
"id": "call_2",
|
|
"type": "function",
|
|
"function": {"name": "search", "arguments": '{"query":"b"}'},
|
|
},
|
|
]
|
|
)
|
|
save_message("s1", "user", "search two things")
|
|
save_message("s1", "assistant", None, tool_calls=tc_json)
|
|
save_message("s1", "tool", "result a", "search", tool_call_id="call_1")
|
|
save_message("s1", "tool", "result b", "search", tool_call_id="call_2")
|
|
msgs = load_messages("s1")
|
|
assert len(msgs) == 4 # user, assistant+2 tool_calls, 2 tool results
|
|
assert len(msgs[1]["tool_calls"]) == 2
|
|
assert msgs[2]["tool_call_id"] == "call_1"
|
|
assert msgs[3]["tool_call_id"] == "call_2"
|
|
|
|
def test_empty_workstream(self, tmp_db):
|
|
assert load_messages("nonexistent") == []
|
|
|
|
|
|
# ── Delete workstream ─────────────────────────────────────────────────
|
|
|
|
|
|
class TestDeleteWorkstream:
|
|
def test_delete_removes_workstream_and_messages(self, tmp_db):
|
|
register_workstream("abc123")
|
|
save_message("abc123", "user", "hello")
|
|
save_message("abc123", "assistant", "hi")
|
|
assert delete_workstream("abc123") is True
|
|
assert list_workstreams_with_history() == []
|
|
assert load_messages("abc123") == []
|
|
|
|
def test_delete_nonexistent(self, tmp_db):
|
|
assert delete_workstream("nonexistent") is False
|
|
|
|
|
|
# ── save_message with tool_call_id ────────────────────────────────────
|
|
|
|
|
|
class TestSaveMessageToolCallId:
|
|
def test_tool_call_id_stored(self, tmp_db):
|
|
save_message("s1", "tool", "output", "bash", tool_call_id="call_xyz")
|
|
engine = get_storage()._engine # noqa: SLF001
|
|
with engine.connect() as conn:
|
|
row = conn.execute(
|
|
sa.text("SELECT tool_call_id FROM conversations WHERE ws_id = 's1'")
|
|
).fetchone()
|
|
assert row[0] == "call_xyz"
|
|
|
|
def test_tool_call_id_none_by_default(self, tmp_db):
|
|
save_message("s1", "user", "hello")
|
|
engine = get_storage()._engine # noqa: SLF001
|
|
with engine.connect() as conn:
|
|
row = conn.execute(
|
|
sa.text("SELECT tool_call_id FROM conversations WHERE ws_id = 's1'")
|
|
).fetchone()
|
|
assert row[0] is None
|
|
|
|
|
|
# ── Workstreams table creation ────────────────────────────────────────
|
|
|
|
|
|
class TestWorkstreamsTable:
|
|
def test_workstreams_table_exists(self, tmp_db):
|
|
engine = get_storage()._engine # noqa: SLF001
|
|
with engine.connect() as conn:
|
|
rows = conn.execute(
|
|
sa.text("SELECT name FROM sqlite_master WHERE type='table' AND name='workstreams'")
|
|
).fetchall()
|
|
assert len(rows) == 1
|
|
|
|
def test_tool_call_id_column_exists(self, tmp_db):
|
|
engine = get_storage()._engine # noqa: SLF001
|
|
with engine.connect() as conn:
|
|
# Should not raise
|
|
conn.execute(sa.text("SELECT tool_call_id FROM conversations LIMIT 0"))
|
|
|
|
|
|
# ── ChatSession.resume ────────────────────────────────────────────────
|
|
|
|
|
|
class TestResumeWorkstream:
|
|
def test_resume_loads_messages(self, tmp_db, mock_openai_client):
|
|
# Set up a workstream with messages in DB
|
|
register_workstream("old_ws_123")
|
|
save_message("old_ws_123", "user", "hello world")
|
|
save_message("old_ws_123", "assistant", "hi there")
|
|
|
|
# Create a new session and resume
|
|
session = ChatSession(
|
|
client=mock_openai_client,
|
|
model="test-model",
|
|
ui=MagicMock(),
|
|
instructions=None,
|
|
temperature=0.5,
|
|
max_tokens=1000,
|
|
tool_timeout=10,
|
|
)
|
|
original_id = session._ws_id
|
|
assert original_id != "old_ws_123"
|
|
|
|
result = session.resume("old_ws_123")
|
|
assert result is True
|
|
assert session._ws_id == "old_ws_123"
|
|
assert len(session.messages) == 2
|
|
assert turn_to_dict(session.messages[0])["content"] == "hello world"
|
|
assert session._title_generated is True
|
|
|
|
def test_resume_nonexistent_returns_false(self, tmp_db, mock_openai_client):
|
|
session = ChatSession(
|
|
client=mock_openai_client,
|
|
model="test-model",
|
|
ui=MagicMock(),
|
|
instructions=None,
|
|
temperature=0.5,
|
|
max_tokens=1000,
|
|
tool_timeout=10,
|
|
)
|
|
assert session.resume("nonexistent") is False
|
|
|
|
def test_workstream_not_registered_until_message(self, tmp_db, mock_openai_client):
|
|
session = ChatSession(
|
|
client=mock_openai_client,
|
|
model="test-model",
|
|
ui=MagicMock(),
|
|
instructions=None,
|
|
temperature=0.5,
|
|
max_tokens=1000,
|
|
tool_timeout=10,
|
|
)
|
|
# Workstream is not auto-registered on init — only on /new or server creation
|
|
assert resolve_workstream(session._ws_id) is None
|
|
assert not any(r[0] == session._ws_id for r in list_workstreams_with_history())
|
|
|
|
|
|
# ── save_message updates workstreams.updated ──────────────────────────
|
|
|
|
|
|
class TestSaveMessageUpdatesWorkstream:
|
|
def test_updated_timestamp_bumped(self, tmp_db):
|
|
register_workstream("s1")
|
|
save_message("s1", "user", "first")
|
|
|
|
import time
|
|
|
|
time.sleep(0.01) # ensure different timestamp
|
|
save_message("s1", "user", "hello")
|
|
|
|
rows = list_workstreams_with_history()
|
|
new_updated = rows[0][4]
|
|
# updated should be same or later (sqlite datetime resolution is seconds,
|
|
# so they may be equal in fast tests — just verify no error)
|
|
assert new_updated is not None
|
|
|
|
|
|
# ── Interrupted workstream repair ─────────────────────────────────────
|
|
|
|
|
|
class TestInterruptedWorkstreamRepair:
|
|
"""load_messages() should strip trailing incomplete tool call turns."""
|
|
|
|
def test_complete_tool_turn_preserved(self, tmp_db):
|
|
"""2 tool_calls + 2 tool results = complete, no stripping."""
|
|
import json
|
|
|
|
tc_json = json.dumps(
|
|
[
|
|
{
|
|
"id": "call_1",
|
|
"type": "function",
|
|
"function": {"name": "bash", "arguments": '{"command":"ls"}'},
|
|
},
|
|
{
|
|
"id": "call_2",
|
|
"type": "function",
|
|
"function": {"name": "bash", "arguments": '{"command":"pwd"}'},
|
|
},
|
|
]
|
|
)
|
|
save_message("s1", "user", "hello")
|
|
save_message("s1", "assistant", None, tool_calls=tc_json)
|
|
save_message("s1", "tool", "file.txt", tool_call_id="call_1")
|
|
save_message("s1", "tool", "/home", tool_call_id="call_2")
|
|
msgs = load_messages("s1")
|
|
assert len(msgs) == 4 # user + assistant(2 calls) + 2 tool results
|
|
|
|
def test_partial_tool_results_stripped(self, tmp_db):
|
|
"""2 tool_calls + 1 tool result = incomplete, strip the turn."""
|
|
import json
|
|
|
|
tc_json = json.dumps(
|
|
[
|
|
{
|
|
"id": "call_1",
|
|
"type": "function",
|
|
"function": {"name": "bash", "arguments": '{"command":"ls"}'},
|
|
},
|
|
{
|
|
"id": "call_2",
|
|
"type": "function",
|
|
"function": {"name": "bash", "arguments": '{"command":"pwd"}'},
|
|
},
|
|
]
|
|
)
|
|
save_message("s1", "user", "hello")
|
|
save_message("s1", "assistant", None, tool_calls=tc_json)
|
|
save_message("s1", "tool", "file.txt", tool_call_id="call_1")
|
|
msgs = load_messages("s1")
|
|
assert len(msgs) == 1 # only user message remains
|
|
assert msgs[0]["role"] == "user"
|
|
|
|
def test_zero_tool_results_stripped(self, tmp_db):
|
|
"""Assistant with tool_calls + 0 results = incomplete, strip the turn."""
|
|
import json
|
|
|
|
tc_json = json.dumps(
|
|
[
|
|
{
|
|
"id": "call_1",
|
|
"type": "function",
|
|
"function": {"name": "bash", "arguments": '{"command":"ls"}'},
|
|
},
|
|
{
|
|
"id": "call_2",
|
|
"type": "function",
|
|
"function": {"name": "bash", "arguments": '{"command":"pwd"}'},
|
|
},
|
|
]
|
|
)
|
|
save_message("s1", "user", "hello")
|
|
save_message("s1", "assistant", "Let me check", tool_calls=tc_json)
|
|
msgs = load_messages("s1")
|
|
assert len(msgs) == 1
|
|
assert msgs[0]["role"] == "user"
|
|
|
|
def test_complete_turn_before_incomplete_preserved(self, tmp_db):
|
|
"""Complete turn followed by incomplete turn: keep complete, strip incomplete."""
|
|
import json
|
|
|
|
tc_json = json.dumps(
|
|
[
|
|
{
|
|
"id": "call_1",
|
|
"type": "function",
|
|
"function": {"name": "bash", "arguments": '{"command":"ls"}'},
|
|
},
|
|
]
|
|
)
|
|
save_message("s1", "user", "first")
|
|
save_message("s1", "assistant", "response")
|
|
save_message("s1", "user", "second")
|
|
save_message("s1", "assistant", None, tool_calls=tc_json)
|
|
msgs = load_messages("s1")
|
|
assert len(msgs) == 3 # user + assistant + user (incomplete turn stripped)
|
|
assert msgs[0]["role"] == "user"
|
|
assert msgs[1]["role"] == "assistant"
|
|
assert msgs[2]["role"] == "user"
|
|
|
|
def test_repair_false_preserves_partial_trailing_turn(self, tmp_db):
|
|
"""``repair=False`` is the display-read contract for ``/history``.
|
|
|
|
The default repair pass strips the trailing
|
|
``assistant(tool_calls)`` when not all tool results are persisted
|
|
— correct for ``session.resume`` (LLM context), wrong for the
|
|
REST display read. A user refreshing the coordinator page mid-
|
|
tool-execution would otherwise lose the entire trailing turn
|
|
from the UI. ``repair=False`` returns the raw persisted state.
|
|
"""
|
|
import json
|
|
|
|
tc_json = json.dumps(
|
|
[
|
|
{
|
|
"id": "call_1",
|
|
"type": "function",
|
|
"function": {"name": "bash", "arguments": '{"command":"ls"}'},
|
|
},
|
|
{
|
|
"id": "call_2",
|
|
"type": "function",
|
|
"function": {"name": "bash", "arguments": '{"command":"pwd"}'},
|
|
},
|
|
]
|
|
)
|
|
save_message("s1", "user", "hello")
|
|
save_message("s1", "assistant", "Checking", tool_calls=tc_json)
|
|
save_message("s1", "tool", "file.txt", tool_call_id="call_1")
|
|
# No call_2 result persisted — mid-execution refresh.
|
|
msgs = get_storage().load_messages("s1", repair=False)
|
|
# All three rows survive — the trailing partial turn is what the
|
|
# operator was actually watching live.
|
|
assert [m["role"] for m in msgs] == ["user", "assistant", "tool"]
|
|
assert msgs[1].get("tool_calls") and len(msgs[1]["tool_calls"]) == 2
|
|
assert msgs[2]["tool_call_id"] == "call_1"
|
|
|
|
def test_repair_false_does_not_synthesize_orphan_results(self, tmp_db):
|
|
"""``repair=False`` must NOT splice synthetic ``"Tool execution
|
|
was cancelled."`` rows for mid-conversation orphans either —
|
|
the operator never saw those rows, and showing them would
|
|
invent UI content that doesn't reflect persisted state.
|
|
"""
|
|
import json
|
|
|
|
tc_json = json.dumps(
|
|
[
|
|
{
|
|
"id": "call_1",
|
|
"type": "function",
|
|
"function": {"name": "bash", "arguments": '{"command":"ls"}'},
|
|
},
|
|
]
|
|
)
|
|
save_message("s1", "user", "first")
|
|
save_message("s1", "assistant", "Working", tool_calls=tc_json)
|
|
# Cancel landed before any tool result — next turn happens.
|
|
save_message("s1", "user", "second")
|
|
save_message("s1", "assistant", "ok")
|
|
msgs = get_storage().load_messages("s1", repair=False)
|
|
roles = [m["role"] for m in msgs]
|
|
# No synthetic tool row spliced after the orphaned tool_calls.
|
|
assert roles == ["user", "assistant", "user", "assistant"]
|
|
assert all(m["role"] != "tool" for m in msgs)
|
|
|
|
|
|
# ── Workstream config persistence ─────────────────────────────────────
|
|
|
|
|
|
class TestWorkstreamConfig:
|
|
def test_save_load_roundtrip(self, tmp_db):
|
|
config = {"temperature": "0.3", "reasoning_effort": "high", "creative_mode": "False"}
|
|
save_workstream_config("s1", config)
|
|
loaded = load_workstream_config("s1")
|
|
assert loaded == config
|
|
|
|
def test_update_existing_key(self, tmp_db):
|
|
save_workstream_config("s1", {"temperature": "0.3"})
|
|
save_workstream_config("s1", {"temperature": "0.7"})
|
|
loaded = load_workstream_config("s1")
|
|
assert loaded["temperature"] == "0.7"
|
|
|
|
def test_missing_workstream_returns_empty(self, tmp_db):
|
|
loaded = load_workstream_config("nonexistent")
|
|
assert loaded == {}
|
|
|
|
def test_delete_workstream_removes_config(self, tmp_db):
|
|
register_workstream("s1")
|
|
save_message("s1", "user", "hi")
|
|
save_workstream_config("s1", {"temperature": "0.5"})
|
|
delete_workstream("s1")
|
|
assert load_workstream_config("s1") == {}
|
|
|
|
def test_resume_restores_config(self, tmp_db):
|
|
"""ChatSession.resume() should restore persisted config."""
|
|
client = MagicMock()
|
|
client.models.list.return_value.data = [MagicMock(id="test-model")]
|
|
ui = MagicMock()
|
|
ui.on_info = MagicMock()
|
|
ui.on_error = MagicMock()
|
|
ui.on_state_change = MagicMock()
|
|
ui.on_rename = MagicMock()
|
|
|
|
# Create a workstream with specific config
|
|
register_workstream("orig")
|
|
save_message("orig", "user", "hello")
|
|
save_message("orig", "assistant", "hi there")
|
|
save_workstream_config(
|
|
"orig",
|
|
{
|
|
"temperature": "0.3",
|
|
"reasoning_effort": "high",
|
|
"max_tokens": "2048",
|
|
"instructions": "be concise",
|
|
"creative_mode": "True",
|
|
},
|
|
)
|
|
|
|
# Create a new session with different defaults, then resume
|
|
session = ChatSession(
|
|
client=client,
|
|
model="test",
|
|
ui=ui,
|
|
instructions=None,
|
|
temperature=0.7,
|
|
max_tokens=4096,
|
|
tool_timeout=30,
|
|
)
|
|
assert session.temperature == 0.7 # default
|
|
result = session.resume("orig")
|
|
assert result is True
|
|
assert session.temperature == 0.3
|
|
assert session.reasoning_effort == "high"
|
|
assert session.max_tokens == 2048
|
|
assert session.instructions == "be concise"
|
|
assert session.creative_mode is True
|
|
|
|
def test_resume_keeps_defaults_when_alias_unresolvable(self, tmp_db):
|
|
"""When the saved alias is empty or no longer in the registry,
|
|
``resume()`` must NOT copy ``saved_model`` onto the constructor's
|
|
default provider. Pairing a removed model name with a default
|
|
provider that doesn't know about it produces a broken session
|
|
whose next API call fails — the exact regression Copilot flagged
|
|
on PR #465. The constructor already resolved a coherent default
|
|
(provider + model + capabilities); resume should leave it intact
|
|
and just log the unreachable saved values."""
|
|
client = MagicMock()
|
|
client.models.list.return_value.data = [MagicMock(id="test-model")]
|
|
ui = MagicMock()
|
|
ui.on_info = MagicMock()
|
|
ui.on_error = MagicMock()
|
|
ui.on_state_change = MagicMock()
|
|
ui.on_rename = MagicMock()
|
|
|
|
register_workstream("model_ws")
|
|
save_message("model_ws", "user", "hello")
|
|
save_message("model_ws", "assistant", "hi")
|
|
# Empty alias + an orphan model name — same shape resume sees
|
|
# when an operator removes an alias from the registry that the
|
|
# workstream was originally pinned to.
|
|
save_workstream_config("model_ws", {"model": "gpt-5", "model_alias": ""})
|
|
|
|
session = ChatSession(
|
|
client=client,
|
|
model="gpt-5-nano",
|
|
ui=ui,
|
|
instructions=None,
|
|
temperature=0.5,
|
|
max_tokens=1000,
|
|
tool_timeout=10,
|
|
)
|
|
assert session.model == "gpt-5-nano"
|
|
result = session.resume("model_ws")
|
|
assert result is True
|
|
# Constructor's coherent default is preserved — saved orphan
|
|
# model name is NOT copied over.
|
|
assert session.model == "gpt-5-nano"
|
|
|
|
def test_init_does_not_clobber_existing_config(self, tmp_db):
|
|
"""ChatSession.__init__ must NOT overwrite existing
|
|
``workstream_config`` keys when constructing for an already-
|
|
persisted ws_id.
|
|
|
|
This is the fix for the rehydrate bug: ``SessionManager.open()``
|
|
builds a ChatSession with the persisted ws_id; the legacy
|
|
``__init__`` unconditionally called ``_save_config()`` which is
|
|
``INSERT OR REPLACE`` per-key — silently resetting model_alias,
|
|
temperature, reasoning_effort, max_tokens, skill, creative_mode,
|
|
and instructions to the constructor defaults *before*
|
|
``resume()`` got a chance to read them back.
|
|
"""
|
|
client = MagicMock()
|
|
client.models.list.return_value.data = [MagicMock(id="test-model")]
|
|
ui = MagicMock()
|
|
ui.on_info = MagicMock()
|
|
ui.on_error = MagicMock()
|
|
ui.on_state_change = MagicMock()
|
|
ui.on_rename = MagicMock()
|
|
|
|
register_workstream("rehydrate_ws")
|
|
save_workstream_config(
|
|
"rehydrate_ws",
|
|
{
|
|
"model": "gpt-5-pro",
|
|
"model_alias": "gpt-5-pro",
|
|
"temperature": "0.2",
|
|
"reasoning_effort": "high",
|
|
"max_tokens": "8192",
|
|
"creative_mode": "True",
|
|
"instructions": "preserve me",
|
|
},
|
|
)
|
|
|
|
ChatSession(
|
|
client=client,
|
|
model="some-default-model",
|
|
ui=ui,
|
|
instructions=None,
|
|
temperature=0.7,
|
|
max_tokens=4096,
|
|
tool_timeout=30,
|
|
reasoning_effort="medium",
|
|
ws_id="rehydrate_ws",
|
|
)
|
|
|
|
loaded = load_workstream_config("rehydrate_ws")
|
|
assert loaded["model"] == "gpt-5-pro"
|
|
assert loaded["model_alias"] == "gpt-5-pro"
|
|
assert loaded["temperature"] == "0.2"
|
|
assert loaded["reasoning_effort"] == "high"
|
|
assert loaded["max_tokens"] == "8192"
|
|
assert loaded["creative_mode"] == "True"
|
|
assert loaded["instructions"] == "preserve me"
|
|
|
|
def test_init_writes_config_on_fresh_create(self, tmp_db):
|
|
"""The opposite half of the contract: when no config row exists
|
|
yet, ``__init__`` must still persist the constructor's values so
|
|
a later resume can find them. This is the path that previously
|
|
worked — the fix must not break it."""
|
|
client = MagicMock()
|
|
client.models.list.return_value.data = [MagicMock(id="test-model")]
|
|
ui = MagicMock()
|
|
ui.on_info = MagicMock()
|
|
ui.on_error = MagicMock()
|
|
ui.on_state_change = MagicMock()
|
|
ui.on_rename = MagicMock()
|
|
|
|
# No save_workstream_config() before ChatSession() — this is
|
|
# the fresh-create path the SessionManager.create() flow takes.
|
|
register_workstream("fresh_ws")
|
|
assert load_workstream_config("fresh_ws") == {}
|
|
|
|
ChatSession(
|
|
client=client,
|
|
model="gpt-5-mini",
|
|
ui=ui,
|
|
instructions="be terse",
|
|
temperature=0.4,
|
|
max_tokens=2048,
|
|
tool_timeout=30,
|
|
reasoning_effort="low",
|
|
ws_id="fresh_ws",
|
|
)
|
|
|
|
loaded = load_workstream_config("fresh_ws")
|
|
assert loaded["model"] == "gpt-5-mini"
|
|
assert loaded["temperature"] == "0.4"
|
|
assert loaded["reasoning_effort"] == "low"
|
|
assert loaded["max_tokens"] == "2048"
|
|
assert loaded["instructions"] == "be terse"
|
|
|
|
|
|
# ── Prune workstreams ─────────────────────────────────────────────────
|
|
|
|
|
|
class TestPruneWorkstreams:
|
|
def test_orphan_removed(self, tmp_db):
|
|
"""Workstream registered with no messages should be pruned."""
|
|
register_workstream("orphan")
|
|
orphans, stale = prune_workstreams()
|
|
assert orphans == 1
|
|
assert list_workstreams_with_history() == []
|
|
|
|
def test_workstream_with_messages_kept(self, tmp_db):
|
|
"""Workstream with messages should not be pruned."""
|
|
register_workstream("active")
|
|
save_message("active", "user", "hello")
|
|
orphans, _stale = prune_workstreams()
|
|
assert orphans == 0
|
|
assert len(list_workstreams_with_history()) == 1
|
|
|
|
def test_stale_unnamed_removed(self, tmp_db):
|
|
"""Old unnamed workstream should be pruned by retention policy."""
|
|
register_workstream("old1")
|
|
save_message("old1", "user", "ancient message")
|
|
# Force the updated timestamp to the past so it looks stale
|
|
engine = get_storage()._engine # noqa: SLF001
|
|
with engine.connect() as conn:
|
|
conn.execute(
|
|
sa.text("UPDATE workstreams SET updated = '2020-01-01' WHERE ws_id = 'old1'")
|
|
)
|
|
conn.commit()
|
|
_orphans, stale = prune_workstreams(retention_days=30)
|
|
assert stale == 1
|
|
|
|
def test_named_workstream_preserved(self, tmp_db):
|
|
"""Workstream with alias should be kept regardless of age."""
|
|
register_workstream("old2")
|
|
set_workstream_alias("old2", "important")
|
|
save_message("old2", "user", "old but named")
|
|
# Force old timestamp
|
|
engine = get_storage()._engine # noqa: SLF001
|
|
with engine.connect() as conn:
|
|
conn.execute(
|
|
sa.text("UPDATE workstreams SET updated = '2020-01-01' WHERE ws_id = 'old2'")
|
|
)
|
|
conn.commit()
|
|
_orphans, stale = prune_workstreams(retention_days=30)
|
|
assert stale == 0
|
|
assert len(list_workstreams_with_history()) == 1
|
|
|
|
def test_fresh_unnamed_preserved(self, tmp_db):
|
|
"""Recent unnamed workstream should not be pruned."""
|
|
register_workstream("fresh")
|
|
save_message("fresh", "user", "just now")
|
|
_orphans, stale = prune_workstreams(retention_days=30)
|
|
assert stale == 0
|
|
assert len(list_workstreams_with_history()) == 1
|
|
|
|
def test_prune_removes_workstream_config(self, tmp_db):
|
|
"""Pruning orphan/stale workstreams should also remove their config rows."""
|
|
register_workstream("orphan_cfg")
|
|
save_workstream_config("orphan_cfg", {"temperature": "0.5"})
|
|
|
|
register_workstream("stale_cfg")
|
|
save_message("stale_cfg", "user", "old")
|
|
save_workstream_config("stale_cfg", {"temperature": "0.9"})
|
|
engine = get_storage()._engine # noqa: SLF001
|
|
with engine.connect() as conn:
|
|
conn.execute(
|
|
sa.text("UPDATE workstreams SET updated = '2020-01-01' WHERE ws_id = 'stale_cfg'")
|
|
)
|
|
conn.commit()
|
|
|
|
# Both should have config before prune
|
|
assert load_workstream_config("orphan_cfg") == {"temperature": "0.5"}
|
|
assert load_workstream_config("stale_cfg") == {"temperature": "0.9"}
|
|
|
|
prune_workstreams(retention_days=30)
|
|
|
|
# Config rows should be cleaned up
|
|
assert load_workstream_config("orphan_cfg") == {}
|
|
assert load_workstream_config("stale_cfg") == {}
|
|
|
|
|
|
# ── Parallel tool exception isolation ────────────────────────────────
|
|
|
|
|
|
class TestParallelToolExceptionIsolation:
|
|
"""Bug #117: one tool raising should not kill the entire batch."""
|
|
|
|
def test_exception_in_one_tool_does_not_kill_batch(self, tmp_db, mock_openai_client):
|
|
from unittest.mock import patch
|
|
|
|
session = ChatSession(
|
|
client=mock_openai_client,
|
|
model="test-model",
|
|
ui=MagicMock(),
|
|
instructions=None,
|
|
temperature=0.5,
|
|
max_tokens=1000,
|
|
tool_timeout=10,
|
|
)
|
|
|
|
def succeed(item):
|
|
return item["call_id"], "ok"
|
|
|
|
def fail(item):
|
|
raise RuntimeError("boom")
|
|
|
|
items = [
|
|
{
|
|
"call_id": "c1",
|
|
"func_name": "bash",
|
|
"execute": succeed,
|
|
"needs_approval": False,
|
|
"header": "test",
|
|
"preview": "",
|
|
},
|
|
{
|
|
"call_id": "c2",
|
|
"func_name": "math",
|
|
"execute": fail,
|
|
"needs_approval": False,
|
|
"header": "test",
|
|
"preview": "",
|
|
},
|
|
]
|
|
|
|
tool_calls = [
|
|
{"id": "c1", "function": {"name": "bash", "arguments": "{}"}},
|
|
{"id": "c2", "function": {"name": "math", "arguments": "{}"}},
|
|
]
|
|
|
|
with (
|
|
patch.object(session, "_prepare_tool", side_effect=items),
|
|
patch.object(session, "_evaluate_intent"),
|
|
patch.object(session, "_emit_state"),
|
|
patch.object(session, "_init_system_messages"),
|
|
patch.object(session, "_check_cancelled"),
|
|
):
|
|
session.ui.approve_tools.return_value = (True, None)
|
|
results, _ = session._execute_tools(tool_calls)
|
|
|
|
assert results[0] == ("c1", "ok")
|
|
assert results[1][0] == "c2"
|
|
assert "Error executing math" in results[1][1]
|
|
assert "boom" in results[1][1]
|
|
|
|
|
|
# ── Web search tool gating ───────────────────────────────────────────
|
|
|
|
|
|
class TestWebSearchGating:
|
|
"""Bug #117: web_search should not be offered without a backend."""
|
|
|
|
def test_web_search_filtered_when_no_backend(self, tmp_db, mock_openai_client):
|
|
from unittest.mock import patch
|
|
|
|
from turnstone.core.providers._protocol import ModelCapabilities
|
|
|
|
session = ChatSession(
|
|
client=mock_openai_client,
|
|
model="local-model",
|
|
ui=MagicMock(),
|
|
instructions=None,
|
|
temperature=0.5,
|
|
max_tokens=1000,
|
|
tool_timeout=10,
|
|
)
|
|
|
|
caps = ModelCapabilities(supports_web_search=False)
|
|
with (
|
|
patch.object(session, "_get_capabilities", return_value=caps),
|
|
patch("turnstone.core.session.get_searxng_url", return_value=None),
|
|
):
|
|
tools = session._get_active_tools()
|
|
|
|
names = [t.get("function", {}).get("name") for t in tools]
|
|
assert "web_search" not in names
|
|
|
|
def test_web_search_kept_when_searxng_configured(self, tmp_db, mock_openai_client):
|
|
from unittest.mock import patch
|
|
|
|
from turnstone.core.providers._protocol import ModelCapabilities
|
|
|
|
session = ChatSession(
|
|
client=mock_openai_client,
|
|
model="local-model",
|
|
ui=MagicMock(),
|
|
instructions=None,
|
|
temperature=0.5,
|
|
max_tokens=1000,
|
|
tool_timeout=10,
|
|
)
|
|
|
|
caps = ModelCapabilities(supports_web_search=False)
|
|
with (
|
|
patch.object(session, "_get_capabilities", return_value=caps),
|
|
patch("turnstone.core.session.get_searxng_url", return_value="http://searxng:8080"),
|
|
):
|
|
tools = session._get_active_tools()
|
|
|
|
names = [t.get("function", {}).get("name") for t in tools]
|
|
assert "web_search" in names
|
|
|
|
def test_web_search_kept_when_native_support(self, tmp_db, mock_openai_client):
|
|
from unittest.mock import patch
|
|
|
|
from turnstone.core.providers._protocol import ModelCapabilities
|
|
|
|
session = ChatSession(
|
|
client=mock_openai_client,
|
|
model="gpt-5-search-api",
|
|
ui=MagicMock(),
|
|
instructions=None,
|
|
temperature=0.5,
|
|
max_tokens=1000,
|
|
tool_timeout=10,
|
|
)
|
|
|
|
caps = ModelCapabilities(supports_web_search=True)
|
|
with (
|
|
patch.object(session, "_get_capabilities", return_value=caps),
|
|
patch("turnstone.core.session.get_searxng_url", return_value=None),
|
|
):
|
|
tools = session._get_active_tools()
|
|
|
|
names = [t.get("function", {}).get("name") for t in tools]
|
|
assert "web_search" in names
|
|
|
|
def test_resolve_search_client_config_store_precedence(self, tmp_db, mock_openai_client):
|
|
"""ConfigStore (admin Settings) is authoritative over the env/config
|
|
fallback: an explicit URL resolves a client; an explicit empty string
|
|
disables web search even when the env var would supply a URL; an unset
|
|
key falls through to env (storage -> toml -> env -> default)."""
|
|
from unittest.mock import patch
|
|
|
|
from turnstone.core.web_search import SearXNGClient
|
|
|
|
class _StubStore:
|
|
def __init__(self, values):
|
|
self._v = values
|
|
|
|
def stored_keys(self):
|
|
return frozenset(self._v)
|
|
|
|
def get(self, key, default=None):
|
|
return self._v.get(key, default)
|
|
|
|
def _session(values):
|
|
return ChatSession(
|
|
client=mock_openai_client,
|
|
model="local-model",
|
|
ui=MagicMock(),
|
|
instructions=None,
|
|
temperature=0.5,
|
|
max_tokens=1000,
|
|
tool_timeout=10,
|
|
config_store=_StubStore(values),
|
|
)
|
|
|
|
# (a) explicit DB URL wins over a None env fallback
|
|
s_db = _session({"tools.searxng_url": "http://db-searx:8080"})
|
|
with patch("turnstone.core.session.get_searxng_url", return_value=None):
|
|
client = s_db._resolve_search_client()
|
|
assert isinstance(client, SearXNGClient)
|
|
assert client._base_url == "http://db-searx:8080"
|
|
|
|
# (b) explicit empty string disables, even though env would supply a URL
|
|
s_off = _session({"tools.searxng_url": ""})
|
|
with patch("turnstone.core.session.get_searxng_url", return_value="http://env-searx:8080"):
|
|
assert s_off._resolve_search_client() is None
|
|
|
|
# (c) unset key falls through to the env/config layer
|
|
s_env = _session({})
|
|
with patch("turnstone.core.session.get_searxng_url", return_value="http://env-searx:8080"):
|
|
client_env = s_env._resolve_search_client()
|
|
assert isinstance(client_env, SearXNGClient)
|
|
assert client_env._base_url == "http://env-searx:8080"
|
|
|
|
|
|
class TestMCPToolGating:
|
|
"""MCP tools should not be offered when no MCP servers provide them."""
|
|
|
|
def test_mcp_tools_filtered_without_mcp_client(self, tmp_db, mock_openai_client):
|
|
"""read_resource and use_prompt excluded when no MCP client."""
|
|
session = ChatSession(
|
|
client=mock_openai_client,
|
|
model="local-model",
|
|
ui=MagicMock(),
|
|
instructions=None,
|
|
temperature=0.5,
|
|
max_tokens=1000,
|
|
tool_timeout=10,
|
|
)
|
|
assert session._mcp_client is None
|
|
|
|
tools = session._get_active_tools()
|
|
names = [t.get("function", {}).get("name") for t in tools]
|
|
assert "read_resource" not in names
|
|
assert "use_prompt" not in names
|
|
|
|
def test_read_resource_filtered_when_no_resources(self, tmp_db, mock_openai_client):
|
|
"""read_resource excluded when MCP client has no resources."""
|
|
mcp_client = MagicMock()
|
|
mcp_client.get_tools.return_value = []
|
|
# Phase 7b: gating uses ``*_count_for_user`` so the test mocks
|
|
# the per-user variant (the property remains for static-only
|
|
# admin paths). Returning 0 / 2 mirrors the prior contract.
|
|
mcp_client.resource_count_for_user.return_value = 0
|
|
mcp_client.prompt_count_for_user.return_value = 2
|
|
|
|
session = ChatSession(
|
|
client=mock_openai_client,
|
|
model="local-model",
|
|
ui=MagicMock(),
|
|
instructions=None,
|
|
temperature=0.5,
|
|
max_tokens=1000,
|
|
tool_timeout=10,
|
|
mcp_client=mcp_client,
|
|
)
|
|
|
|
tools = session._get_active_tools()
|
|
names = [t.get("function", {}).get("name") for t in tools]
|
|
assert "read_resource" not in names
|
|
assert "use_prompt" in names
|
|
|
|
def test_use_prompt_filtered_when_no_prompts(self, tmp_db, mock_openai_client):
|
|
"""use_prompt excluded when MCP client has no prompts."""
|
|
mcp_client = MagicMock()
|
|
mcp_client.get_tools.return_value = []
|
|
mcp_client.resource_count_for_user.return_value = 3
|
|
mcp_client.prompt_count_for_user.return_value = 0
|
|
|
|
session = ChatSession(
|
|
client=mock_openai_client,
|
|
model="local-model",
|
|
ui=MagicMock(),
|
|
instructions=None,
|
|
temperature=0.5,
|
|
max_tokens=1000,
|
|
tool_timeout=10,
|
|
mcp_client=mcp_client,
|
|
)
|
|
|
|
tools = session._get_active_tools()
|
|
names = [t.get("function", {}).get("name") for t in tools]
|
|
assert "use_prompt" not in names
|
|
assert "read_resource" in names
|
|
|
|
def test_mcp_tools_kept_when_servers_have_both(self, tmp_db, mock_openai_client):
|
|
"""Both tools present when MCP client has resources and prompts."""
|
|
mcp_client = MagicMock()
|
|
mcp_client.get_tools.return_value = []
|
|
mcp_client.resource_count_for_user.return_value = 1
|
|
mcp_client.prompt_count_for_user.return_value = 1
|
|
|
|
session = ChatSession(
|
|
client=mock_openai_client,
|
|
model="local-model",
|
|
ui=MagicMock(),
|
|
instructions=None,
|
|
temperature=0.5,
|
|
max_tokens=1000,
|
|
tool_timeout=10,
|
|
mcp_client=mcp_client,
|
|
)
|
|
|
|
tools = session._get_active_tools()
|
|
names = [t.get("function", {}).get("name") for t in tools]
|
|
assert "read_resource" in names
|
|
assert "use_prompt" in names
|
|
|
|
def test_mcp_tools_filtered_with_tool_search_active(self, tmp_db, mock_openai_client):
|
|
"""Gating applies even when tool_search is active (client-side path)."""
|
|
mcp_client = MagicMock()
|
|
mcp_client.get_tools.return_value = []
|
|
mcp_client.resource_count_for_user.return_value = 0
|
|
mcp_client.prompt_count_for_user.return_value = 0
|
|
|
|
session = ChatSession(
|
|
client=mock_openai_client,
|
|
model="local-model",
|
|
ui=MagicMock(),
|
|
instructions=None,
|
|
temperature=0.5,
|
|
max_tokens=1000,
|
|
tool_timeout=10,
|
|
mcp_client=mcp_client,
|
|
tool_search="on",
|
|
)
|
|
assert session._tool_search is not None
|
|
|
|
tools = session._get_active_tools()
|
|
names = [t.get("function", {}).get("name") for t in tools]
|
|
assert "read_resource" not in names
|
|
assert "use_prompt" not in names
|
|
|
|
def test_mcp_tools_filtered_with_native_tool_search(self, tmp_db, mock_openai_client):
|
|
"""Gating applies when provider handles tool search natively."""
|
|
from unittest.mock import patch
|
|
|
|
from turnstone.core.providers._protocol import ModelCapabilities
|
|
|
|
mcp_client = MagicMock()
|
|
mcp_client.get_tools.return_value = []
|
|
mcp_client.resource_count_for_user.return_value = 0
|
|
mcp_client.prompt_count_for_user.return_value = 0
|
|
|
|
session = ChatSession(
|
|
client=mock_openai_client,
|
|
model="local-model",
|
|
ui=MagicMock(),
|
|
instructions=None,
|
|
temperature=0.5,
|
|
max_tokens=1000,
|
|
tool_timeout=10,
|
|
mcp_client=mcp_client,
|
|
tool_search="on",
|
|
)
|
|
|
|
caps = ModelCapabilities(supports_tool_search=True)
|
|
with patch.object(session, "_get_capabilities", return_value=caps):
|
|
tools = session._get_active_tools()
|
|
|
|
names = [t.get("function", {}).get("name") for t in tools]
|
|
assert "read_resource" not in names
|
|
assert "use_prompt" not in names
|
|
|
|
def test_pool_only_user_keeps_read_resource_and_use_prompt(self, tmp_db, mock_openai_client):
|
|
"""Phase 7b canary: a pool-only user (static catalog empty) still
|
|
sees ``read_resource`` and ``use_prompt`` because the gating
|
|
consults ``*_count_for_user`` (scope decision 0.2).
|
|
|
|
Drives ``resource_count = prompt_count = 0`` (the static-only
|
|
properties are zero) but ``*_count_for_user(uid) > 0`` because
|
|
the user has pool entries; the tools must remain visible.
|
|
"""
|
|
mcp_client = MagicMock()
|
|
mcp_client.get_tools.return_value = []
|
|
# Static catalog is empty; admin-style legacy properties say 0.
|
|
mcp_client.resource_count = 0
|
|
mcp_client.prompt_count = 0
|
|
# Per-user variant reports the user's pool entries.
|
|
mcp_client.resource_count_for_user.return_value = 2
|
|
mcp_client.prompt_count_for_user.return_value = 1
|
|
|
|
session = ChatSession(
|
|
client=mock_openai_client,
|
|
model="local-model",
|
|
ui=MagicMock(),
|
|
instructions=None,
|
|
temperature=0.5,
|
|
max_tokens=1000,
|
|
tool_timeout=10,
|
|
mcp_client=mcp_client,
|
|
user_id="pool-only-user",
|
|
)
|
|
|
|
tools = session._get_active_tools()
|
|
names = [t.get("function", {}).get("name") for t in tools]
|
|
assert "read_resource" in names
|
|
assert "use_prompt" in names
|
|
# Verify the per-user gate was actually consulted with the
|
|
# session's ``user_id`` (sanity-check on the wiring).
|
|
mcp_client.resource_count_for_user.assert_any_call("pool-only-user")
|
|
mcp_client.prompt_count_for_user.assert_any_call("pool-only-user")
|
|
|
|
|
|
class TestMCPActingUserBinding:
|
|
"""Per-user MCP credentials follow the acting user on shared workstreams.
|
|
|
|
The workstream owner is the fallback identity; an authenticated send
|
|
rebinds credential resolution (dispatch + catalogs + listeners) to the
|
|
sender. Prepared tool items pin the identity at prepare time so a
|
|
pending approval can't execute under a later sender's credentials.
|
|
"""
|
|
|
|
def _make(self, mock_openai_client, owner="alice"):
|
|
mcp_client = MagicMock()
|
|
mcp_client.get_tools.return_value = []
|
|
mcp_client.call_tool_sync.return_value = "ok"
|
|
session = ChatSession(
|
|
client=mock_openai_client,
|
|
model="local-model",
|
|
ui=MagicMock(),
|
|
instructions=None,
|
|
temperature=0.5,
|
|
max_tokens=1000,
|
|
tool_timeout=10,
|
|
mcp_client=mcp_client,
|
|
user_id=owner,
|
|
)
|
|
# Capture instead of persisting — same stub idiom as
|
|
# test_session_mcp_dispatch_error.
|
|
session._report_tool_result = MagicMock() # type: ignore[method-assign]
|
|
return session, mcp_client
|
|
|
|
def test_effective_identity_defaults_to_owner(self, tmp_db, mock_openai_client):
|
|
session, mcp_client = self._make(mock_openai_client)
|
|
assert session._mcp_effective_user_id == "alice"
|
|
item = session._prepare_mcp_tool("c1", "mcp__srv__tool", {})
|
|
session._exec_mcp_tool(item)
|
|
assert mcp_client.call_tool_sync.call_args.kwargs["user_id"] == "alice"
|
|
|
|
def test_bind_rebinds_dispatch_catalog_listeners_and_prime(self, tmp_db, mock_openai_client):
|
|
session, mcp_client = self._make(mock_openai_client)
|
|
mcp_client.reset_mock()
|
|
|
|
session.bind_acting_user("bob")
|
|
|
|
# Dispatch identity follows the acting user.
|
|
item = session._prepare_mcp_tool("c1", "mcp__srv__tool", {})
|
|
session._exec_mcp_tool(item)
|
|
assert mcp_client.call_tool_sync.call_args.kwargs["user_id"] == "bob"
|
|
# Listener registrations swapped from owner to acting user for
|
|
# all three catalog kinds — identity is the (user_id, callback)
|
|
# pair, so the remove must name the OLD uid and the add the new.
|
|
mcp_client.remove_listener.assert_called_once_with(session._mcp_refresh_cb, user_id="alice")
|
|
mcp_client.add_listener.assert_called_once_with(session._mcp_refresh_cb, user_id="bob")
|
|
mcp_client.remove_resource_listener.assert_called_once_with(
|
|
session._mcp_resource_cb, user_id="alice"
|
|
)
|
|
mcp_client.add_resource_listener.assert_called_once_with(
|
|
session._mcp_resource_cb, user_id="bob"
|
|
)
|
|
mcp_client.remove_prompt_listener.assert_called_once_with(
|
|
session._mcp_prompt_cb, user_id="alice"
|
|
)
|
|
mcp_client.add_prompt_listener.assert_called_once_with(
|
|
session._mcp_prompt_cb, user_id="bob"
|
|
)
|
|
# The acting user's oauth_user pools are warmed so their tools
|
|
# surface without a manual reconnect.
|
|
mcp_client.prime_user_pools.assert_called_once_with("bob")
|
|
# Merged tool list rebuilt under the new identity.
|
|
mcp_client.get_tools.assert_any_call(user_id="bob")
|
|
|
|
def test_prepared_item_pins_identity_across_rebind(self, tmp_db, mock_openai_client):
|
|
session, mcp_client = self._make(mock_openai_client)
|
|
session.bind_acting_user("bob")
|
|
item = session._prepare_mcp_tool("c1", "mcp__srv__tool", {})
|
|
# A different user takes over the session while the item is
|
|
# pending approval — execution must stay under the requester.
|
|
session.bind_acting_user("carol")
|
|
session._exec_mcp_tool(item)
|
|
assert mcp_client.call_tool_sync.call_args.kwargs["user_id"] == "bob"
|
|
|
|
def test_resource_and_prompt_items_pin_identity(self, tmp_db, mock_openai_client):
|
|
session, mcp_client = self._make(mock_openai_client)
|
|
session.bind_acting_user("bob")
|
|
res_item = session._prepare_read_resource("c1", {"uri": "res://x"})
|
|
mcp_client.is_mcp_prompt.return_value = True
|
|
prompt_item = session._prepare_use_prompt("c2", {"name": "p"})
|
|
session.bind_acting_user("carol")
|
|
assert res_item["mcp_user_id"] == "bob"
|
|
assert prompt_item["mcp_user_id"] == "bob"
|
|
# And the prompt-existence gate consults the CURRENT effective
|
|
# identity (carol) for new preparations.
|
|
session._prepare_use_prompt("c3", {"name": "p"})
|
|
assert mcp_client.is_mcp_prompt.call_args.kwargs["user_id"] == "carol"
|
|
|
|
def test_bind_noops_on_empty_and_same_user(self, tmp_db, mock_openai_client):
|
|
session, mcp_client = self._make(mock_openai_client)
|
|
mcp_client.reset_mock()
|
|
session.bind_acting_user("")
|
|
session.bind_acting_user("alice") # same as owner
|
|
mcp_client.remove_listener.assert_not_called()
|
|
mcp_client.add_listener.assert_not_called()
|
|
mcp_client.prime_user_pools.assert_not_called()
|
|
assert session._mcp_effective_user_id == "alice"
|
|
|
|
def test_send_kwarg_binds_before_turn_starts(self, tmp_db, mock_openai_client):
|
|
import pytest
|
|
|
|
session, _mcp_client = self._make(mock_openai_client)
|
|
|
|
class _SentinelError(Exception):
|
|
pass
|
|
|
|
# ``bind_acting_user`` runs before ``_refresh_model_from_registry``
|
|
# at the top of send() — abort there to prove the ordering without
|
|
# driving the full agent loop.
|
|
session._refresh_model_from_registry = MagicMock( # type: ignore[method-assign]
|
|
side_effect=_SentinelError
|
|
)
|
|
with pytest.raises(_SentinelError):
|
|
session.send("hi", acting_user_id="bob")
|
|
assert session._acting_user_id == "bob"
|
|
|
|
def test_close_removes_listeners_under_rebound_identity(self, tmp_db, mock_openai_client):
|
|
session, mcp_client = self._make(mock_openai_client)
|
|
session.bind_acting_user("bob")
|
|
refresh_cb = session._mcp_refresh_cb
|
|
mcp_client.reset_mock()
|
|
session.close()
|
|
mcp_client.remove_listener.assert_called_once_with(refresh_cb, user_id="bob")
|
|
|
|
def test_bind_without_mcp_client_only_records(self, tmp_db, mock_openai_client):
|
|
session = ChatSession(
|
|
client=mock_openai_client,
|
|
model="local-model",
|
|
ui=MagicMock(),
|
|
instructions=None,
|
|
temperature=0.5,
|
|
max_tokens=1000,
|
|
tool_timeout=10,
|
|
user_id="alice",
|
|
)
|
|
session.bind_acting_user("bob")
|
|
assert session._mcp_effective_user_id == "bob"
|