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:
co-authored by
Claude Opus 4.6
parent
777a09c806
commit
2508f36ebf
@@ -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()
|
||||
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user