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
|
@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()
|
||||||
|
|
||||||
|
|||||||
@@ -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,6 +112,8 @@ 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 = _provision_sessions.get(provision_id)
|
||||||
|
if session:
|
||||||
session["status"] = "completed"
|
session["status"] = "completed"
|
||||||
session["username"] = result.login_name
|
session["username"] = result.login_name
|
||||||
logger.info(
|
logger.info(
|
||||||
@@ -116,6 +123,8 @@ async def _poll_and_store(provision_id: str) -> None:
|
|||||||
return
|
return
|
||||||
|
|
||||||
if result.status == "expired":
|
if result.status == "expired":
|
||||||
|
session = _provision_sessions.get(provision_id)
|
||||||
|
if session:
|
||||||
session["status"] = "expired"
|
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,6 +134,8 @@ async def _poll_and_store(provision_id: str) -> None:
|
|||||||
await anyio.sleep(2)
|
await anyio.sleep(2)
|
||||||
|
|
||||||
# Timed out
|
# Timed out
|
||||||
|
session = _provision_sessions.get(provision_id)
|
||||||
|
if session:
|
||||||
session["status"] = "expired"
|
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")
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user