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
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()
+30 -10
View File
@@ -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")
+192 -7
View File
@@ -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