diff --git a/docs/source/reference/rest-api/websockets.md b/docs/source/reference/rest-api/websockets.md index 81b54ecdac..548e2323cb 100644 --- a/docs/source/reference/rest-api/websockets.md +++ b/docs/source/reference/rest-api/websockets.md @@ -63,6 +63,19 @@ to the client. - `error`: Error information object with `code` (string, see Error types), `message` (string), and `details` (string) - `schema_version`: schema version - `OPTIONAL` +## Reconnecting an Active Conversation + +To resume an active workflow, reconnect with the workflow's `conversation_id` query parameter and an identity credential that resolves to the same user who started it. The server restores state only when both values match. A connection with another identity, or no identity, does not receive the existing workflow state or its pending Human-in-the-Loop prompt. + +The server accepts the following identity credentials: + +- A `nat-session` cookie or `?session=` query parameter. +- A JWT Bearer token. +- An API key supplied as a Bearer token or `X-API-Key` header. +- HTTP Basic credentials. + +Clients can also send a JWT, API key, or Basic credentials through an `auth_message`. Send the message before expecting restoration; restoration occurs after authentication succeeds. + ## Auth Message This message allows clients to authenticate over a WebSocket connection when header-based or cookie-based authentication is not feasible (e.g., browser WebSocket APIs that do not support custom headers). diff --git a/packages/nvidia_nat_core/src/nat/front_ends/fastapi/fastapi_front_end_plugin_worker.py b/packages/nvidia_nat_core/src/nat/front_ends/fastapi/fastapi_front_end_plugin_worker.py index dddb65162f..7ee9b9eab3 100644 --- a/packages/nvidia_nat_core/src/nat/front_ends/fastapi/fastapi_front_end_plugin_worker.py +++ b/packages/nvidia_nat_core/src/nat/front_ends/fastapi/fastapi_front_end_plugin_worker.py @@ -209,8 +209,8 @@ def __init__(self, config: Config): self._outstanding_flows: dict[str, FlowState] = {} self._outstanding_flows_lock = asyncio.Lock() - # Conversation handlers for WebSocket reconnection support - self._conversation_handlers: dict[str, WebSocketMessageHandler] = {} + # Conversation handlers for identity-bound WebSocket reconnection support + self._conversation_handlers: dict[tuple[str, str], WebSocketMessageHandler] = {} # Track session managers for each route self._session_managers: list[SessionManager] = [] @@ -228,17 +228,17 @@ def __init__(self, config: Config): remove_flow_cb=self._remove_flow, ) - def get_conversation_handler(self, conversation_id: str) -> "WebSocketMessageHandler | None": - """Get a conversation handler for reconnection support.""" - return self._conversation_handlers.get(conversation_id) + def get_conversation_handler(self, user_id: str, conversation_id: str) -> "WebSocketMessageHandler | None": + """Get the conversation handler owned by a user.""" + return self._conversation_handlers.get((user_id, conversation_id)) - def set_conversation_handler(self, conversation_id: str, handler: "WebSocketMessageHandler") -> None: - """Register a conversation handler for reconnection support.""" - self._conversation_handlers[conversation_id] = handler + def set_conversation_handler(self, user_id: str, conversation_id: str, handler: "WebSocketMessageHandler") -> None: + """Register a conversation handler under its user and conversation IDs.""" + self._conversation_handlers[(user_id, conversation_id)] = handler - def remove_conversation_handler(self, conversation_id: str) -> None: - """Remove a conversation handler when workflow completes.""" - self._conversation_handlers.pop(conversation_id, None) + def remove_conversation_handler(self, user_id: str, conversation_id: str) -> None: + """Remove a user's conversation handler when its workflow completes.""" + self._conversation_handlers.pop((user_id, conversation_id), None) async def initialize_evaluators(self, config: Config): """Initialize and store evaluators from config for single-item evaluation.""" diff --git a/packages/nvidia_nat_core/src/nat/front_ends/fastapi/message_handler.py b/packages/nvidia_nat_core/src/nat/front_ends/fastapi/message_handler.py index 104fa3ccd9..76a0dc6de1 100644 --- a/packages/nvidia_nat_core/src/nat/front_ends/fastapi/message_handler.py +++ b/packages/nvidia_nat_core/src/nat/front_ends/fastapi/message_handler.py @@ -103,6 +103,7 @@ def __init__(self, self._user_interaction: UserInteraction | None = None self._pending_observability_trace: ResponseObservabilityTrace | None = None self._user_id: str | None = None + self._restoration_attempted: bool = False self._flow_handler: FlowHandlerBase | None = None @@ -127,16 +128,20 @@ def _initialize_workflow_request(self, message: WebSocketUserMessage) -> None: self._workflow_schema_type = message.schema_type self._conversation_id = message.conversation_id self._user_message_payload: dict[str, Any] = message.model_dump() - if self._conversation_id: - self._worker.set_conversation_handler(self._conversation_id, self) + if self._user_id and self._conversation_id: + self._worker.set_conversation_handler(self._user_id, self._conversation_id, self) async def _restore_execution_state(self) -> None: """Restore execution state on reconnection by swapping handler state.""" + if self._restoration_attempted or not self._user_id: + return + + self._restoration_attempted = True conversation_id = self._socket.query_params.get("conversation_id") if not conversation_id: return - disconnected_handler = self._worker.get_conversation_handler(conversation_id) + disconnected_handler = self._worker.get_conversation_handler(self._user_id, conversation_id) if not disconnected_handler: return @@ -170,6 +175,9 @@ async def _restore_execution_state(self) -> None: async def __aenter__(self) -> "WebSocketMessageHandler": await self._socket.accept() + user_info = UserManager.extract_user_from_connection(self._socket) + if user_info is not None: + self._user_id = user_info.get_user_id() await self._restore_execution_state() return self @@ -273,9 +281,11 @@ async def _process_auth_message(self, message: WebSocketAuthMessage) -> None: self._flow_handler.set_oauth_mode(message.payload.mode) return + identity_resolved = False try: user_info: UserInfo = UserManager._from_auth_payload(message.payload) self._user_id = user_info.get_user_id() + identity_resolved = True response: WebSocketAuthResponseMessage = WebSocketAuthResponseMessage( status=AuthMessageStatus.SUCCESS, user_id=self._user_id, @@ -290,6 +300,8 @@ async def _process_auth_message(self, message: WebSocketAuthMessage) -> None: ), ) await self._socket.send_json(response.model_dump()) + if identity_resolved: + await self._restore_execution_state() async def _process_websocket_user_interaction_response_message( self, user_content: WebSocketUserInteractionResponseMessage) -> TextContent: @@ -334,13 +346,14 @@ async def process_workflow_request(self, user_message_as_validated_type: WebSock self._running_workflow_task = None _conversation_id = self._conversation_id + _user_id = self._user_id def _done_callback(_task: asyncio.Task): if self._running_workflow_task is _task: self._running_workflow_task = None - if self._running_workflow_task is None and _conversation_id and \ - self._worker.get_conversation_handler(_conversation_id) is self: - self._worker.remove_conversation_handler(_conversation_id) + if self._running_workflow_task is None and _user_id and _conversation_id and \ + self._worker.get_conversation_handler(_user_id, _conversation_id) is self: + self._worker.remove_conversation_handler(_user_id, _conversation_id) # Only the *_STREAM schemas stream; others aggregate a single result. Streaming a # non-streaming schema converts chunks to the single output schema and raises. diff --git a/packages/nvidia_nat_core/tests/nat/front_ends/fastapi/test_message_handler.py b/packages/nvidia_nat_core/tests/nat/front_ends/fastapi/test_message_handler.py index 60c5dcf2df..bee8546794 100644 --- a/packages/nvidia_nat_core/tests/nat/front_ends/fastapi/test_message_handler.py +++ b/packages/nvidia_nat_core/tests/nat/front_ends/fastapi/test_message_handler.py @@ -16,9 +16,12 @@ import asyncio from unittest.mock import AsyncMock from unittest.mock import MagicMock +from unittest.mock import patch from starlette.websockets import WebSocketDisconnect +from nat.data_models.api_server import ApiKeyAuthPayload +from nat.data_models.api_server import AuthMessageStatus from nat.data_models.api_server import OAuthMode from nat.data_models.api_server import OAuthModePreferencePayload from nat.data_models.api_server import WebSocketAuthMessage @@ -49,6 +52,71 @@ def _make_message_handler() -> tuple[WebSocketMessageHandler, AsyncMock, WebSock return handler, socket, flow_handler +async def test_context_manager_resolves_connection_identity_before_restoration(): + """Connection credentials establish the owner before reconnection is attempted.""" + handler, socket, _ = _make_message_handler() + user_info = MagicMock() + user_info.get_user_id.return_value = "user-a" + restore = AsyncMock() + handler._restore_execution_state = restore + + with patch("nat.front_ends.fastapi.message_handler.UserManager.extract_user_from_connection", + return_value=user_info): + await handler.__aenter__() + + socket.accept.assert_awaited_once() + assert handler._user_id == "user-a" + restore.assert_awaited_once() + + +async def test_anonymous_connection_does_not_attempt_conversation_lookup(): + """A conversation ID cannot restore state without a resolved user identity.""" + handler, socket, _ = _make_message_handler() + socket.query_params = {"conversation_id": "conversation-a"} + + await handler._restore_execution_state() + + handler._worker.get_conversation_handler.assert_not_called() + + +async def test_successful_auth_message_attempts_owned_restoration_once(): + """Delayed authentication can restore once and cannot retry the lookup.""" + handler, socket, _ = _make_message_handler() + socket.query_params = {"conversation_id": "conversation-a"} + handler._worker.get_conversation_handler.return_value = None + user_info = MagicMock() + user_info.get_user_id.return_value = "user-a" + msg = WebSocketAuthMessage( + type=WebSocketMessageType.AUTH_MESSAGE, + payload=ApiKeyAuthPayload(method="api_key", token="test-api-key"), + ) + + with patch("nat.front_ends.fastapi.message_handler.UserManager._from_auth_payload", return_value=user_info): + await handler._process_auth_message(msg) + await handler._process_auth_message(msg) + + handler._worker.get_conversation_handler.assert_called_once_with("user-a", "conversation-a") + + +async def test_failed_auth_message_does_not_attempt_restoration(): + """A failed identity resolution cannot trigger conversation restoration.""" + handler, socket, _ = _make_message_handler() + restore = AsyncMock() + handler._restore_execution_state = restore + msg = WebSocketAuthMessage( + type=WebSocketMessageType.AUTH_MESSAGE, + payload=ApiKeyAuthPayload(method="api_key", token="test-api-key"), + ) + + with patch("nat.front_ends.fastapi.message_handler.UserManager._from_auth_payload", + side_effect=ValueError("invalid credential")): + await handler._process_auth_message(msg) + + restore.assert_not_awaited() + response = socket.send_json.await_args.args[0] + assert response["status"] == AuthMessageStatus.ERROR + + async def test_process_auth_message_sets_oauth_mode_and_sends_no_response(): """An oauth_mode_preference payload updates the flow handler's mode and emits no auth response.""" handler, socket, flow_handler = _make_message_handler() diff --git a/packages/nvidia_nat_core/tests/nat/server/test_unified_api_server.py b/packages/nvidia_nat_core/tests/nat/server/test_unified_api_server.py index dc82c72b41..124292d743 100644 --- a/packages/nvidia_nat_core/tests/nat/server/test_unified_api_server.py +++ b/packages/nvidia_nat_core/tests/nat/server/test_unified_api_server.py @@ -1246,6 +1246,7 @@ async def test_restore_execution_state_sends_prompt_with_remaining_timeout(): ) handler.create_websocket_message = AsyncMock() handler._conversation_id = "conv1" + handler._user_id = "user-a" future: asyncio.Future = asyncio.get_running_loop().create_future() prompt_content = HumanPromptText(text="Confirm?", required=True, placeholder="y", timeout=10) @@ -1264,12 +1265,33 @@ async def test_restore_execution_state_sends_prompt_with_remaining_timeout(): with patch("nat.front_ends.fastapi.message_handler.time.monotonic", return_value=3.0): await handler._restore_execution_state() + mock_worker.get_conversation_handler.assert_called_once_with("user-a", "conv1") handler.create_websocket_message.assert_called_once() call_kwargs = handler.create_websocket_message.call_args[1] sent_content = call_kwargs["data_model"] assert sent_content.timeout == 7 +def test_conversation_handler_registry_isolates_users_and_owner_cleanup(): + """Equal conversation IDs remain isolated across users and clean up independently.""" + worker = FastApiFrontEndPluginWorker.__new__(FastApiFrontEndPluginWorker) + worker._conversation_handlers = {} + user_a_handler = MagicMock() + user_b_handler = MagicMock() + + worker.set_conversation_handler("user-a", "shared-conversation", user_a_handler) + worker.set_conversation_handler("user-b", "shared-conversation", user_b_handler) + + assert worker.get_conversation_handler("user-a", "shared-conversation") is user_a_handler + assert worker.get_conversation_handler("user-b", "shared-conversation") is user_b_handler + assert worker.get_conversation_handler("anonymous", "shared-conversation") is None + + worker.remove_conversation_handler("user-a", "shared-conversation") + + assert worker.get_conversation_handler("user-a", "shared-conversation") is None + assert worker.get_conversation_handler("user-b", "shared-conversation") is user_b_handler + + async def test_process_workflow_request_cancels_in_flight_task(): """A new workflow request cancels any in-flight task before creating a replacement.""" mock_socket = AsyncMock()