From 2508f36ebfdbafc321ca05c79adaa27a47ce4ee7 Mon Sep 17 00:00:00 2001 From: Chris Coutinho Date: Mon, 30 Mar 2026 09:29:14 +0200 Subject: [PATCH] =?UTF-8?q?fix:=20address=20PR=20review=20round=202=20?= =?UTF-8?q?=E2=80=94=20expiry=20checks,=20race=20guards,=20poll=20tests?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Log warning if /app mount not found when sharing poll task group - Add docstring explaining unconditional task group creation - Check session expires_at in provision_status to catch stale sessions - Guard _poll_and_store status writes against cleanup-while-polling race - Use "error" status (not "expired") when app_password is missing - Remove hardcoded "Astrolabe Background Sync" user_agent string - Fix async mock pattern (new_callable=AsyncMock) in test - Add autouse fixture to clear _provision_sessions between tests - Add _poll_and_store unit tests: completed, expired, error, cleanup - Document all status values in provision_status docstring Co-Authored-By: Claude Opus 4.6 (1M context) --- nextcloud_mcp_server/app.py | 14 +- nextcloud_mcp_server/auth/provision_routes.py | 40 +++- tests/unit/test_provision_routes.py | 199 +++++++++++++++++- 3 files changed, 235 insertions(+), 18 deletions(-) diff --git a/nextcloud_mcp_server/app.py b/nextcloud_mcp_server/app.py index f755e137..e0d60554 100644 --- a/nextcloud_mcp_server/app.py +++ b/nextcloud_mcp_server/app.py @@ -1413,16 +1413,28 @@ def get_app(transport: str = "streamable-http", enabled_apps: list[str] | None = @asynccontextmanager async def _maybe_login_flow_cleanup(app: Starlette): - """Start Login Flow cleanup task and provision poll task group.""" + """Start Login Flow cleanup task and provision poll task group. + + The task group is always created (even when Login Flow cleanup is + disabled) because provision routes use it to spawn background poll + tasks via ``browser_app.state.poll_task_group``. + """ async with anyio.create_task_group() as tg: if settings.enable_login_flow: tg.start_soon(_login_flow_cleanup_loop) # Share task group with provision routes for background polling + found_app_mount = False for route in app.routes: if isinstance(route, Mount) and route.path == "/app": browser_app = cast(Starlette, route.app) browser_app.state.poll_task_group = tg + found_app_mount = True break + if not found_app_mount: + logger.warning( + "Could not find /app mount to share poll task group; " + "web provisioning will return 500" + ) yield tg.cancel_scope.cancel() diff --git a/nextcloud_mcp_server/auth/provision_routes.py b/nextcloud_mcp_server/auth/provision_routes.py index 179b6901..530057c0 100644 --- a/nextcloud_mcp_server/auth/provision_routes.py +++ b/nextcloud_mcp_server/auth/provision_routes.py @@ -64,7 +64,8 @@ async def _poll_and_store(provision_id: str) -> None: settings = get_settings() nextcloud_host = settings.nextcloud_host if not nextcloud_host: - session["status"] = "expired" + if provision_id in _provision_sessions: + session["status"] = "error" return flow_client = LoginFlowV2Client( @@ -96,7 +97,11 @@ async def _poll_and_store(provision_id: str) -> None: storage = await get_shared_storage() effective_user_id = user_id or result.login_name or "unknown" if not result.app_password: - session["status"] = "expired" + # Re-fetch session to avoid writing to orphaned dict if + # _cleanup_expired_sessions removed it while we were polling + session = _provision_sessions.get(provision_id) + if session: + session["status"] = "error" logger.error( f"Login Flow v2 completed but no app_password (provision_id={provision_id})" ) @@ -107,8 +112,10 @@ async def _poll_and_store(provision_id: str) -> None: scopes=None, # All scopes username=result.login_name, ) - session["status"] = "completed" - session["username"] = result.login_name + session = _provision_sessions.get(provision_id) + if session: + session["status"] = "completed" + session["username"] = result.login_name logger.info( f"Login Flow v2 web provision completed for user {effective_user_id} " f"(provision_id={provision_id})" @@ -116,7 +123,9 @@ async def _poll_and_store(provision_id: str) -> None: return if result.status == "expired": - session["status"] = "expired" + session = _provision_sessions.get(provision_id) + if session: + session["status"] = "expired" logger.warning( f"Login Flow v2 web provision expired (provision_id={provision_id})" ) @@ -125,7 +134,9 @@ async def _poll_and_store(provision_id: str) -> None: await anyio.sleep(2) # Timed out - session["status"] = "expired" + session = _provision_sessions.get(provision_id) + if session: + session["status"] = "expired" logger.warning( f"Login Flow v2 web provision timed out (provision_id={provision_id})" ) @@ -177,9 +188,7 @@ async def provision_page(request: Request) -> RedirectResponse | HTMLResponse: nextcloud_host=nextcloud_host, verify_ssl=get_nextcloud_ssl_verify(), ) - init_response = await flow_client.initiate( - user_agent="Astrolabe Background Sync" - ) + init_response = await flow_client.initiate() except Exception as e: logger.error(f"Failed to initiate Login Flow v2 for web provision: {e}") return HTMLResponse( @@ -234,7 +243,12 @@ async def provision_status(request: Request) -> JSONResponse: GET /app/provision/status?id=... - Returns JSON: {"status": "pending"|"completed"|"expired", "username": "..."} + Returns JSON with status field: + - ``"pending"`` — flow in progress, poll again + - ``"completed"`` — app password stored, includes ``"username"`` + - ``"expired"`` — flow timed out or was rejected by Nextcloud + - ``"error"`` — flow completed but server-side error (e.g. missing app password) + - ``"not_found"`` — unknown or already-consumed session (404) """ provision_id = request.query_params.get("id", "") @@ -248,6 +262,12 @@ async def provision_status(request: Request) -> JSONResponse: status_code=404, ) + # Detect sessions that outlived their TTL (e.g. no new provision + # requests triggered _cleanup_expired_sessions) + if session["expires_at"] < time.time(): + _provision_sessions.pop(provision_id, None) + return JSONResponse({"status": "expired"}, status_code=404) + response: dict = {"status": session["status"]} if session["status"] == "completed": response["username"] = session.get("username") diff --git a/tests/unit/test_provision_routes.py b/tests/unit/test_provision_routes.py index 5cd67b16..a14b5ee0 100644 --- a/tests/unit/test_provision_routes.py +++ b/tests/unit/test_provision_routes.py @@ -1,6 +1,7 @@ """Unit tests for web-based Login Flow v2 provisioning routes. -Tests validation, HTML escaping, URL rewriting, and route handlers. +Tests validation, HTML escaping, URL rewriting, route handlers, and +background polling logic. """ import time @@ -8,7 +9,9 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest +from nextcloud_mcp_server.auth.login_flow import LoginFlowPollResult from nextcloud_mcp_server.auth.provision_routes import ( + _poll_and_store, _provision_sessions, _render_error, _validate_redirect_uri, @@ -19,6 +22,14 @@ from nextcloud_mcp_server.auth.provision_routes import ( pytestmark = pytest.mark.unit +@pytest.fixture(autouse=True) +def _clear_provision_sessions(): + """Ensure _provision_sessions is empty before and after each test.""" + _provision_sessions.clear() + yield + _provision_sessions.clear() + + # ── _validate_redirect_uri tests ───────────────────────────────────────── @@ -106,12 +117,9 @@ async def test_provision_status_pending(): "status": "pending", "expires_at": time.time() + 600, } - try: - request = _make_request({"id": provision_id}) - response = await provision_status(request) - assert response.status_code == 200 - finally: - _provision_sessions.pop(provision_id, None) + request = _make_request({"id": provision_id}) + response = await provision_status(request) + assert response.status_code == 200 async def test_provision_status_completed_cleans_up(): @@ -129,6 +137,19 @@ async def test_provision_status_completed_cleans_up(): assert provision_id not in _provision_sessions +async def test_provision_status_expired_by_ttl(): + """Session past its TTL is reported as expired and cleaned up.""" + provision_id = "test-expired-ttl" + _provision_sessions[provision_id] = { + "status": "pending", + "expires_at": time.time() - 1, # Already expired + } + request = _make_request({"id": provision_id}) + response = await provision_status(request) + assert response.status_code == 404 + assert provision_id not in _provision_sessions + + # ── provision_page tests ───────────────────────────────────────────────── @@ -162,9 +183,173 @@ async def test_provision_page_skips_if_already_provisioned(): with patch( "nextcloud_mcp_server.auth.provision_routes.get_shared_storage", + new_callable=AsyncMock, return_value=mock_storage, ): response = await provision_page(request) assert response.status_code == 307 # RedirectResponse default assert response.headers["location"] == "https://app.example.com/settings" + + +# ── _poll_and_store tests ──────────────────────────────────────────────── + + +def _create_poll_session(provision_id: str) -> dict: + """Create a minimal provision session for polling tests.""" + session = { + "status": "pending", + "poll_endpoint": "https://cloud.example.com/login/v2/poll", + "poll_token": "secret-token", + "user_id": "alice", + "created_at": time.time(), + "expires_at": time.time() + 1200, + } + _provision_sessions[provision_id] = session + return session + + +async def test_poll_and_store_completed(): + """Successful poll stores app password and sets status to completed.""" + provision_id = "test-poll-completed" + _create_poll_session(provision_id) + + mock_poll_result = LoginFlowPollResult( + status="completed", + server="https://cloud.example.com", + login_name="alice", + app_password="aaaaa-bbbbb-ccccc-ddddd-eeeee", + ) + + mock_flow_client = AsyncMock() + mock_flow_client.poll.return_value = mock_poll_result + + mock_storage = AsyncMock() + + mock_settings = MagicMock() + mock_settings.nextcloud_host = "https://cloud.example.com" + + with ( + patch( + "nextcloud_mcp_server.auth.provision_routes.get_settings", + return_value=mock_settings, + ), + patch( + "nextcloud_mcp_server.auth.provision_routes.get_nextcloud_ssl_verify", + return_value=False, + ), + patch( + "nextcloud_mcp_server.auth.provision_routes.LoginFlowV2Client", + return_value=mock_flow_client, + ), + patch( + "nextcloud_mcp_server.auth.provision_routes.get_shared_storage", + new_callable=AsyncMock, + return_value=mock_storage, + ), + ): + await _poll_and_store(provision_id) + + assert _provision_sessions[provision_id]["status"] == "completed" + assert _provision_sessions[provision_id]["username"] == "alice" + mock_storage.store_app_password_with_scopes.assert_called_once_with( + user_id="alice", + app_password="aaaaa-bbbbb-ccccc-ddddd-eeeee", + scopes=None, + username="alice", + ) + + +async def test_poll_and_store_expired(): + """Expired poll result sets session status to expired.""" + provision_id = "test-poll-expired" + _create_poll_session(provision_id) + + mock_poll_result = LoginFlowPollResult(status="expired") + + mock_flow_client = AsyncMock() + mock_flow_client.poll.return_value = mock_poll_result + + mock_settings = MagicMock() + mock_settings.nextcloud_host = "https://cloud.example.com" + + with ( + patch( + "nextcloud_mcp_server.auth.provision_routes.get_settings", + return_value=mock_settings, + ), + patch( + "nextcloud_mcp_server.auth.provision_routes.get_nextcloud_ssl_verify", + return_value=False, + ), + patch( + "nextcloud_mcp_server.auth.provision_routes.LoginFlowV2Client", + return_value=mock_flow_client, + ), + ): + await _poll_and_store(provision_id) + + assert _provision_sessions[provision_id]["status"] == "expired" + + +async def test_poll_and_store_missing_app_password(): + """Completed poll with no app_password sets status to error.""" + provision_id = "test-poll-no-password" + _create_poll_session(provision_id) + + mock_poll_result = LoginFlowPollResult( + status="completed", + server="https://cloud.example.com", + login_name="alice", + app_password=None, # Missing + ) + + mock_flow_client = AsyncMock() + mock_flow_client.poll.return_value = mock_poll_result + + mock_settings = MagicMock() + mock_settings.nextcloud_host = "https://cloud.example.com" + + mock_storage = AsyncMock() + + with ( + patch( + "nextcloud_mcp_server.auth.provision_routes.get_settings", + return_value=mock_settings, + ), + patch( + "nextcloud_mcp_server.auth.provision_routes.get_nextcloud_ssl_verify", + return_value=False, + ), + patch( + "nextcloud_mcp_server.auth.provision_routes.LoginFlowV2Client", + return_value=mock_flow_client, + ), + patch( + "nextcloud_mcp_server.auth.provision_routes.get_shared_storage", + new_callable=AsyncMock, + return_value=mock_storage, + ), + ): + await _poll_and_store(provision_id) + + assert _provision_sessions[provision_id]["status"] == "error" + mock_storage.store_app_password_with_scopes.assert_not_called() + + +async def test_poll_and_store_session_cleaned_up(): + """Poll exits early if session was cleaned up externally.""" + provision_id = "test-poll-cleaned" + # Don't create session — simulate it being cleaned up before poll starts + + mock_settings = MagicMock() + mock_settings.nextcloud_host = "https://cloud.example.com" + + with patch( + "nextcloud_mcp_server.auth.provision_routes.get_settings", + return_value=mock_settings, + ): + # Should return immediately without error + await _poll_and_store(provision_id) + + assert provision_id not in _provision_sessions