fix: address PR review round 2 — expiry checks, race guards, poll tests

- 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) <noreply@anthropic.com>
This commit is contained in:
Chris Coutinho
2026-03-30 09:29:14 +02:00
co-authored by Claude Opus 4.6
parent 777a09c806
commit 2508f36ebf
3 changed files with 235 additions and 18 deletions
+13 -1
View File
@@ -1413,16 +1413,28 @@ def get_app(transport: str = "streamable-http", enabled_apps: list[str] | None =
@asynccontextmanager @asynccontextmanager
async def _maybe_login_flow_cleanup(app: Starlette): 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: async with anyio.create_task_group() as tg:
if settings.enable_login_flow: if settings.enable_login_flow:
tg.start_soon(_login_flow_cleanup_loop) tg.start_soon(_login_flow_cleanup_loop)
# Share task group with provision routes for background polling # Share task group with provision routes for background polling
found_app_mount = False
for route in app.routes: for route in app.routes:
if isinstance(route, Mount) and route.path == "/app": if isinstance(route, Mount) and route.path == "/app":
browser_app = cast(Starlette, route.app) browser_app = cast(Starlette, route.app)
browser_app.state.poll_task_group = tg browser_app.state.poll_task_group = tg
found_app_mount = True
break break
if not found_app_mount:
logger.warning(
"Could not find /app mount to share poll task group; "
"web provisioning will return 500"
)
yield yield
tg.cancel_scope.cancel() tg.cancel_scope.cancel()
+30 -10
View File
@@ -64,7 +64,8 @@ async def _poll_and_store(provision_id: str) -> None:
settings = get_settings() settings = get_settings()
nextcloud_host = settings.nextcloud_host nextcloud_host = settings.nextcloud_host
if not nextcloud_host: if not nextcloud_host:
session["status"] = "expired" if provision_id in _provision_sessions:
session["status"] = "error"
return return
flow_client = LoginFlowV2Client( flow_client = LoginFlowV2Client(
@@ -96,7 +97,11 @@ async def _poll_and_store(provision_id: str) -> None:
storage = await get_shared_storage() storage = await get_shared_storage()
effective_user_id = user_id or result.login_name or "unknown" effective_user_id = user_id or result.login_name or "unknown"
if not result.app_password: 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( logger.error(
f"Login Flow v2 completed but no app_password (provision_id={provision_id})" 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 scopes=None, # All scopes
username=result.login_name, username=result.login_name,
) )
session["status"] = "completed" session = _provision_sessions.get(provision_id)
session["username"] = result.login_name if session:
session["status"] = "completed"
session["username"] = result.login_name
logger.info( logger.info(
f"Login Flow v2 web provision completed for user {effective_user_id} " f"Login Flow v2 web provision completed for user {effective_user_id} "
f"(provision_id={provision_id})" f"(provision_id={provision_id})"
@@ -116,7 +123,9 @@ async def _poll_and_store(provision_id: str) -> None:
return return
if result.status == "expired": if result.status == "expired":
session["status"] = "expired" session = _provision_sessions.get(provision_id)
if session:
session["status"] = "expired"
logger.warning( logger.warning(
f"Login Flow v2 web provision expired (provision_id={provision_id})" 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) await anyio.sleep(2)
# Timed out # Timed out
session["status"] = "expired" session = _provision_sessions.get(provision_id)
if session:
session["status"] = "expired"
logger.warning( logger.warning(
f"Login Flow v2 web provision timed out (provision_id={provision_id})" 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, nextcloud_host=nextcloud_host,
verify_ssl=get_nextcloud_ssl_verify(), verify_ssl=get_nextcloud_ssl_verify(),
) )
init_response = await flow_client.initiate( init_response = await flow_client.initiate()
user_agent="Astrolabe Background Sync"
)
except Exception as e: except Exception as e:
logger.error(f"Failed to initiate Login Flow v2 for web provision: {e}") logger.error(f"Failed to initiate Login Flow v2 for web provision: {e}")
return HTMLResponse( return HTMLResponse(
@@ -234,7 +243,12 @@ async def provision_status(request: Request) -> JSONResponse:
GET /app/provision/status?id=... 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", "") provision_id = request.query_params.get("id", "")
@@ -248,6 +262,12 @@ async def provision_status(request: Request) -> JSONResponse:
status_code=404, 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"]} response: dict = {"status": session["status"]}
if session["status"] == "completed": if session["status"] == "completed":
response["username"] = session.get("username") response["username"] = session.get("username")
+192 -7
View File
@@ -1,6 +1,7 @@
"""Unit tests for web-based Login Flow v2 provisioning routes. """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 import time
@@ -8,7 +9,9 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest import pytest
from nextcloud_mcp_server.auth.login_flow import LoginFlowPollResult
from nextcloud_mcp_server.auth.provision_routes import ( from nextcloud_mcp_server.auth.provision_routes import (
_poll_and_store,
_provision_sessions, _provision_sessions,
_render_error, _render_error,
_validate_redirect_uri, _validate_redirect_uri,
@@ -19,6 +22,14 @@ from nextcloud_mcp_server.auth.provision_routes import (
pytestmark = pytest.mark.unit 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 ───────────────────────────────────────── # ── _validate_redirect_uri tests ─────────────────────────────────────────
@@ -106,12 +117,9 @@ async def test_provision_status_pending():
"status": "pending", "status": "pending",
"expires_at": time.time() + 600, "expires_at": time.time() + 600,
} }
try: request = _make_request({"id": provision_id})
request = _make_request({"id": provision_id}) response = await provision_status(request)
response = await provision_status(request) assert response.status_code == 200
assert response.status_code == 200
finally:
_provision_sessions.pop(provision_id, None)
async def test_provision_status_completed_cleans_up(): 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 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 ───────────────────────────────────────────────── # ── provision_page tests ─────────────────────────────────────────────────
@@ -162,9 +183,173 @@ async def test_provision_page_skips_if_already_provisioned():
with patch( with patch(
"nextcloud_mcp_server.auth.provision_routes.get_shared_storage", "nextcloud_mcp_server.auth.provision_routes.get_shared_storage",
new_callable=AsyncMock,
return_value=mock_storage, return_value=mock_storage,
): ):
response = await provision_page(request) response = await provision_page(request)
assert response.status_code == 307 # RedirectResponse default assert response.status_code == 307 # RedirectResponse default
assert response.headers["location"] == "https://app.example.com/settings" 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