diff --git a/tests/test_sessions.py b/tests/test_sessions.py index c97ad1d2..5bd99af8 100644 --- a/tests/test_sessions.py +++ b/tests/test_sessions.py @@ -1191,3 +1191,147 @@ class TestMCPToolGating: # 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" diff --git a/turnstone/core/session.py b/turnstone/core/session.py index df159122..3e06fc2c 100644 --- a/turnstone/core/session.py +++ b/turnstone/core/session.py @@ -1174,6 +1174,18 @@ class ChatSession: # Listener identity, catalog merge, and dispatch all consume # this through ``_mcp_user_id``. self._mcp_user_id: str | None = user_id or None + # Acting user for per-user MCP credential resolution on SHARED + # workstreams: the authenticated principal who most recently + # initiated a turn (bound by ``bind_acting_user`` from the send + # path). Empty until the first authenticated send, and empty + # forever on CLI / eval / scheduled sessions — every consumer + # goes through ``_mcp_effective_user_id``, which falls back to + # the session owner. Rebinding also swaps the user-scoped MCP + # listeners; ``_mcp_listener_user_id`` tracks the identity the + # current registrations were made under (listener identity is + # the ``(user_id, callback)`` pair). + self._acting_user_id: str = "" + self._mcp_listener_user_id: str | None = user_id or None self._username = username self._client_type = client_type # Whether the user is online to complete an in-flight OAuth @@ -2210,11 +2222,12 @@ class ChatSession: # is fixed at COORDINATOR_TOOLS. Ignore MCP server changes. if self._kind == WorkstreamKind.COORDINATOR: return - # Phase 7: pass session-bound user_id so the merged tool list - # includes this user's pool catalog. The static path is included - # by ``get_tools`` regardless; ``user_id=None`` would silently - # drop pool tools that the LLM is allowed to call. - mcp_tools = self._mcp_client.get_tools(user_id=self._mcp_user_id) + # Pass the effective user_id (acting user on shared workstreams, + # owner otherwise) so the merged tool list includes that user's + # pool catalog. The static path is included by ``get_tools`` + # regardless; ``user_id=None`` would silently drop pool tools + # that the LLM is allowed to call. + mcp_tools = self._mcp_client.get_tools(user_id=self._mcp_effective_user_id) self._tools = merge_mcp_tools(INTERACTIVE_TOOLS, mcp_tools) self._task_tools = merge_mcp_tools(TASK_AGENT_TOOLS, mcp_tools) self._render_agent_tool_descriptions() @@ -2447,20 +2460,26 @@ class ChatSession: if self._mcp_client and self._mcp_refresh_cb: # ``user_id`` MUST match the value used at registration — # the listener identity is ``(user_id, callback)``, not - # callback alone. ``self._user_id`` is set once in - # ``__init__`` and never mutated, so identity is stable. - self._mcp_client.remove_listener(self._mcp_refresh_cb, user_id=self._mcp_user_id) + # callback alone. ``bind_acting_user`` may have re-scoped + # the registrations since construction, so the tracked + # ``_mcp_listener_user_id`` (not ``_mcp_user_id``) is the + # registration identity. + self._mcp_client.remove_listener( + self._mcp_refresh_cb, user_id=self._mcp_listener_user_id + ) self._mcp_refresh_cb = None if self._mcp_client and self._mcp_resource_cb: # ``user_id`` MUST mirror the value passed at registration — # the listener identity is ``(user_id, callback)`` and an # unscoped removal would leave the registration in place. self._mcp_client.remove_resource_listener( - self._mcp_resource_cb, user_id=self._mcp_user_id + self._mcp_resource_cb, user_id=self._mcp_listener_user_id ) self._mcp_resource_cb = None if self._mcp_client and self._mcp_prompt_cb: - self._mcp_client.remove_prompt_listener(self._mcp_prompt_cb, user_id=self._mcp_user_id) + self._mcp_client.remove_prompt_listener( + self._mcp_prompt_cb, user_id=self._mcp_listener_user_id + ) self._mcp_prompt_cb = None if self._watch_runner: self._watch_runner.remove_dispatch_fn(self._ws_id) @@ -3113,9 +3132,10 @@ class ChatSession: ) # MCP resource catalog (lets the model know what's available for read_resource) if self._mcp_client: - # Per-user merge: pool entries for ``self._mcp_user_id`` are - # included; other users' pool resources are not. - all_resources = self._mcp_client.get_resources(user_id=self._mcp_user_id) + # Per-user merge: pool entries for the effective user (acting + # user on shared workstreams, owner otherwise) are included; + # other users' pool resources are not. + all_resources = self._mcp_client.get_resources(user_id=self._mcp_effective_user_id) concrete = [r for r in all_resources if not r.get("template")] templates = [r for r in all_resources if r.get("template")] if concrete or templates: @@ -3140,9 +3160,10 @@ class ChatSession: dev_parts.append("\n".join(lines)) # MCP prompt catalog (lets the model know what's available for use_prompt) if self._mcp_client: - # Per-user merge: pool entries for ``self._mcp_user_id`` are - # included; other users' pool prompts are not. - prompts = self._mcp_client.get_prompts(user_id=self._mcp_user_id) + # Per-user merge: pool entries for the effective user (acting + # user on shared workstreams, owner otherwise) are included; + # other users' pool prompts are not. + prompts = self._mcp_client.get_prompts(user_id=self._mcp_effective_user_id) if prompts: lines = [""] for p in prompts[:30]: @@ -3931,9 +3952,13 @@ class ChatSession: # connected. Per-user variants (scope decision 0.2) keep the # tool visible for a pool-only user even when the static catalog # is empty. - if not self._mcp_client or not self._mcp_client.resource_count_for_user(self._mcp_user_id): + if not self._mcp_client or not self._mcp_client.resource_count_for_user( + self._mcp_effective_user_id + ): tools = _without_tool(tools, "read_resource") - if not self._mcp_client or not self._mcp_client.prompt_count_for_user(self._mcp_user_id): + if not self._mcp_client or not self._mcp_client.prompt_count_for_user( + self._mcp_effective_user_id + ): tools = _without_tool(tools, "use_prompt") return tools @@ -4463,6 +4488,81 @@ class ChatSession: # -- Main generation loop ------------------------------------------------ + @property + def _mcp_effective_user_id(self) -> str | None: + """Identity for per-user MCP (oauth_user) credential resolution. + + The acting user — the authenticated principal who last initiated + a turn on this session — when one is bound; otherwise the session + owner. On a shared workstream this makes MCP tool calls run under + the credentials (and catalog) of whoever is actually driving, + rather than whoever created the workstream. Falls back to the + owner for CLI / eval / scheduled / internal turns, preserving the + pre-existing single-user behaviour. + """ + return self._acting_user_id or self._mcp_user_id + + def bind_acting_user(self, user_id: str) -> None: + """Bind the authenticated initiator of the current turn. + + Called from the HTTP send path with the caller's authenticated + user id. No-ops when ``user_id`` is empty (unauthenticated lanes + keep the owner fallback) or unchanged. On a genuine change this + re-scopes the session's MCP view to the new acting user: + + - swaps the user-scoped tool/resource/prompt listeners so pool + catalog changes for the acting user reach this session + (listener identity is the ``(user_id, callback)`` pair); + - fire-and-forget primes the acting user's oauth_user pools so + their tools surface without a manual reconnect; + - rebuilds the merged tool list and catalog-dependent state via + the same callbacks a pool notification would fire. + + The binding is sticky — it persists until the next authenticated + send — so wake nudges and auto-resume continuations keep running + under the user whose turn they continue. It intentionally does + NOT rebind mid-turn: queued interjections fold into the current + turn under the initiator's identity, and prepared tool items pin + the identity at prepare time (see ``_prepare_mcp_tool``). + """ + if not user_id or user_id == (self._acting_user_id or self._user_id): + self._acting_user_id = self._acting_user_id or user_id + return + self._acting_user_id = user_id + mcp = self._mcp_client + if not mcp or self._kind == WorkstreamKind.COORDINATOR: + return + old_listener_uid = self._mcp_listener_user_id + new_listener_uid: str | None = self._mcp_effective_user_id + if new_listener_uid != old_listener_uid: + if self._mcp_refresh_cb: + mcp.remove_listener(self._mcp_refresh_cb, user_id=old_listener_uid) + mcp.add_listener(self._mcp_refresh_cb, user_id=new_listener_uid) + if self._mcp_resource_cb: + mcp.remove_resource_listener(self._mcp_resource_cb, user_id=old_listener_uid) + mcp.add_resource_listener(self._mcp_resource_cb, user_id=new_listener_uid) + if self._mcp_prompt_cb: + mcp.remove_prompt_listener(self._mcp_prompt_cb, user_id=old_listener_uid) + mcp.add_prompt_listener(self._mcp_prompt_cb, user_id=new_listener_uid) + self._mcp_listener_user_id = new_listener_uid + if new_listener_uid and hasattr(mcp, "prime_user_pools"): + try: + mcp.prime_user_pools(new_listener_uid) + except Exception: + log.debug( + "mcp prime_user_pools scheduling failed user=%s", + new_listener_uid, + exc_info=True, + ) + # Rebuild the merged tool list and resource/prompt-dependent + # state under the new identity NOW — the prime above completes + # asynchronously and only notifies on catalog changes, while + # already-warm pool entries for this user produce no + # notification at all. + self._on_mcp_tools_changed() + self._on_mcp_resources_changed() + self._on_mcp_prompts_changed() + def send( self, user_input: str, @@ -4470,6 +4570,7 @@ class ChatSession: send_id: str | None = None, *, from_wake: bool = False, + acting_user_id: str | None = None, ) -> None: """Send user input and handle the response loop (including tool calls). @@ -4483,7 +4584,15 @@ class ChatSession: ``send_id`` is an end-to-end tracking token only; it no longer gates a DB reservation (the upload buffer is the pending store, and the bytes in ``attachments`` were already drained/peeked from it by the caller). + + ``acting_user_id`` is the authenticated caller who initiated this + turn (HTTP send / retry paths). It rebinds per-user MCP credential + resolution to that user for this and subsequent turns — see + :meth:`bind_acting_user`. ``None`` (internal callers: wake nudges, + auto-resume, CLI) leaves the current binding untouched. """ + if acting_user_id is not None: + self.bind_acting_user(acting_user_id) self._refresh_model_from_registry() # Token budget approval gate if self._budget_exhausted: @@ -7640,13 +7749,13 @@ class ChatSession: } preparer = preparers.get(func_name) if not preparer: - # Check if this is an MCP tool. Phase 7: pass session-bound - # ``user_id`` so per-user pool tools become reachable here — - # without this kwarg the gate stays static-only and pool - # dispatch is structurally unreachable from - # ``ChatSession._prepare_tool`` (RFC §3, invariant 8). + # Check if this is an MCP tool. Pass the effective ``user_id`` + # (acting user on shared workstreams) so per-user pool tools + # become reachable here — without this kwarg the gate stays + # static-only and pool dispatch is structurally unreachable + # from ``ChatSession._prepare_tool`` (RFC §3, invariant 8). if self._mcp_client and self._mcp_client.is_mcp_tool( - func_name, user_id=self._mcp_user_id + func_name, user_id=self._mcp_effective_user_id ): return self._prepare_mcp_tool(call_id, func_name, args) self.ui.on_error(f"Model called unknown tool: {func_name!r}") @@ -7655,7 +7764,7 @@ class ChatSession: available.extend( sorted( t["function"]["name"] - for t in self._mcp_client.get_tools(user_id=self._mcp_user_id) + for t in self._mcp_client.get_tools(user_id=self._mcp_effective_user_id) ) ) return { @@ -11783,6 +11892,10 @@ class ChatSession: "execute": self._exec_mcp_tool, "mcp_func_name": func_name, "mcp_args": args, + # Pin the credential identity at prepare time: an item that + # sits pending approval must execute under the user whose + # turn requested it, not whoever binds the session later. + "mcp_user_id": self._mcp_effective_user_id, } def _exec_mcp_tool(self, item: dict[str, Any]) -> tuple[str, str]: @@ -11799,7 +11912,7 @@ class ChatSession: output = self._mcp_client.call_tool_sync( func_name, args, - user_id=self._mcp_user_id, + user_id=item.get("mcp_user_id", self._mcp_effective_user_id), timeout=self.tool_timeout, is_interactive_for_consent=self._is_interactive_for_consent, ) @@ -11873,6 +11986,8 @@ class ChatSession: "approval_label": f"mcp_resource__{self._normalize_resource_uri(uri)}", "execute": self._exec_read_resource, "resource_uri": uri, + # Pinned at prepare time — see _prepare_mcp_tool. + "mcp_user_id": self._mcp_effective_user_id, } def _exec_read_resource(self, item: dict[str, Any]) -> tuple[str, str]: @@ -11891,7 +12006,7 @@ class ChatSession: # static path runs byte-identical (invariant 1). output = self._mcp_client.read_resource_sync( uri, - user_id=self._mcp_user_id, + user_id=item.get("mcp_user_id", self._mcp_effective_user_id), timeout=self.tool_timeout, is_interactive_for_consent=self._is_interactive_for_consent, ) @@ -11934,7 +12049,7 @@ class ChatSession: "needs_approval": False, "error": "No MCP servers configured", } - if not self._mcp_client.is_mcp_prompt(name, user_id=self._mcp_user_id): + if not self._mcp_client.is_mcp_prompt(name, user_id=self._mcp_effective_user_id): return { "call_id": call_id, "func_name": "use_prompt", @@ -11968,6 +12083,8 @@ class ChatSession: "execute": self._exec_use_prompt, "prompt_name": name, "prompt_arguments": arguments, + # Pinned at prepare time — see _prepare_mcp_tool. + "mcp_user_id": self._mcp_effective_user_id, } def _exec_use_prompt(self, item: dict[str, Any]) -> tuple[str, str]: @@ -11988,7 +12105,7 @@ class ChatSession: messages = self._mcp_client.get_prompt_sync( name, arguments or None, - user_id=self._mcp_user_id, + user_id=item.get("mcp_user_id", self._mcp_effective_user_id), timeout=self.tool_timeout, is_interactive_for_consent=self._is_interactive_for_consent, ) @@ -14419,12 +14536,12 @@ class ChatSession: elif arg and arg.split()[0] == "refresh": self._handle_mcp_refresh(arg) else: - # Phase 7 + 7b: pass session-bound user_id so the /mcp + # Phase 7 + 7b: pass the effective user_id so the /mcp # listing surfaces this user's pool tools, resources, # and prompts alongside the static catalog. - tools = self._mcp_client.get_tools(user_id=self._mcp_user_id) - resources = self._mcp_client.get_resources(user_id=self._mcp_user_id) - prompts = self._mcp_client.get_prompts(user_id=self._mcp_user_id) + tools = self._mcp_client.get_tools(user_id=self._mcp_effective_user_id) + resources = self._mcp_client.get_resources(user_id=self._mcp_effective_user_id) + prompts = self._mcp_client.get_prompts(user_id=self._mcp_effective_user_id) mcp_lines = [] if tools: mcp_lines.append(f"MCP tools ({len(tools)}):") diff --git a/turnstone/core/session_routes.py b/turnstone/core/session_routes.py index b5cf89f3..a4e56ace 100644 --- a/turnstone/core/session_routes.py +++ b/turnstone/core/session_routes.py @@ -1540,6 +1540,16 @@ def make_retry_handler( ui._enqueue({"type": "busy_error", "message": "Cannot retry while processing."}) return JSONResponse({"status": "busy"}) + # A retry is a fresh turn initiated by the authenticated caller — + # rebind per-user MCP credential resolution to them before the + # re-send dispatches (the per-kind ``dispatch_retry`` closure + # calls ``send()`` without identity kwargs). + from turnstone.core.web_helpers import auth_user_id + + acting_uid = auth_user_id(request) + if acting_uid: + session.bind_acting_user(acting_uid) + retry_msg = session.retry() if hasattr(ui, "_enqueue"): @@ -3620,7 +3630,7 @@ def make_send_handler(cfg: SessionEndpointConfig) -> Handler: from turnstone.core import session_worker from turnstone.core.session import AttachmentsNotQueueableError, GenerationCancelled - from turnstone.core.web_helpers import read_json_or_400 + from turnstone.core.web_helpers import auth_user_id, read_json_or_400 async def send(request: Request) -> Response: if cfg.permission_gate is not None: @@ -3653,6 +3663,14 @@ def make_send_handler(cfg: SessionEndpointConfig) -> Handler: if ui is None: return JSONResponse({"error": "session UI not available"}, status_code=409) + # The authenticated sender: threaded into the fresh-turn dispatch + # below so per-user MCP credentials follow whoever is actually + # driving a shared workstream. Deliberately NOT applied on the + # live-worker queue path — an interjection folds into the current + # turn under the initiator's identity (no mid-turn credential + # switch); the next fresh turn rebinds. + acting_uid = auth_user_id(request) + # ----- Attachment resolution (from the per-node upload buffer) ----- send_id = "" requested_ids: list[str] = [] @@ -3758,6 +3776,8 @@ def make_send_handler(cfg: SessionEndpointConfig) -> Handler: kwargs["attachments"] = resolved_atts if send_id: kwargs["send_id"] = send_id + if acting_uid: + kwargs["acting_user_id"] = acting_uid session.send(message, **kwargs) except GenerationCancelled: # Safety net — send() normally handles this internally.