Merge pull request #669 from cbcoutinho/feat/login-flow-v2-web-provision
feat: add web-based Login Flow v2 provisioning endpoint
This commit is contained in:
+4
-2
@@ -3,9 +3,11 @@ FROM docker.io/library/python:3.12-slim-trixie@sha256:f3fa41d74a768c2fce8016b98c
|
|||||||
COPY --from=ghcr.io/astral-sh/uv:0.10.12@sha256:72ab0aeb448090480ccabb99fb5f52b0dc3c71923bffb5e2e26517a1c27b7fec /uv /uvx /bin/
|
COPY --from=ghcr.io/astral-sh/uv:0.10.12@sha256:72ab0aeb448090480ccabb99fb5f52b0dc3c71923bffb5e2e26517a1c27b7fec /uv /uvx /bin/
|
||||||
|
|
||||||
# Install dependencies
|
# Install dependencies
|
||||||
# 1. git (required for caldav dependency from git)
|
# 1. curl (required for container healthcheck probes)
|
||||||
# 2. sqlite for development with token db
|
# 2. git (required for caldav dependency from git)
|
||||||
|
# 3. sqlite for development with token db
|
||||||
RUN apt update && apt install --no-install-recommends --no-install-suggests -y \
|
RUN apt update && apt install --no-install-recommends --no-install-suggests -y \
|
||||||
|
curl \
|
||||||
git \
|
git \
|
||||||
tesseract-ocr \
|
tesseract-ocr \
|
||||||
sqlite3 && apt clean
|
sqlite3 && apt clean
|
||||||
|
|||||||
+50
-16
@@ -73,6 +73,10 @@ from nextcloud_mcp_server.auth.oauth_routes import (
|
|||||||
oauth_register_proxy,
|
oauth_register_proxy,
|
||||||
oauth_token_endpoint,
|
oauth_token_endpoint,
|
||||||
)
|
)
|
||||||
|
from nextcloud_mcp_server.auth.provision_routes import (
|
||||||
|
provision_page,
|
||||||
|
provision_status,
|
||||||
|
)
|
||||||
from nextcloud_mcp_server.auth.session_backend import SessionAuthBackend
|
from nextcloud_mcp_server.auth.session_backend import SessionAuthBackend
|
||||||
from nextcloud_mcp_server.auth.storage import RefreshTokenStorage, get_shared_storage
|
from nextcloud_mcp_server.auth.storage import RefreshTokenStorage, get_shared_storage
|
||||||
from nextcloud_mcp_server.auth.token_broker import TokenBrokerService
|
from nextcloud_mcp_server.auth.token_broker import TokenBrokerService
|
||||||
@@ -1394,6 +1398,9 @@ def get_app(transport: str = "streamable-http", enabled_apps: list[str] | None =
|
|||||||
from nextcloud_mcp_server.auth.oauth_routes import ( # noqa: PLC0415
|
from nextcloud_mcp_server.auth.oauth_routes import ( # noqa: PLC0415
|
||||||
_cleanup_expired_proxy_codes,
|
_cleanup_expired_proxy_codes,
|
||||||
)
|
)
|
||||||
|
from nextcloud_mcp_server.auth.provision_routes import ( # noqa: PLC0415
|
||||||
|
_cleanup_expired_sessions as _cleanup_expired_provision_sessions,
|
||||||
|
)
|
||||||
|
|
||||||
while True:
|
while True:
|
||||||
try:
|
try:
|
||||||
@@ -1403,27 +1410,45 @@ def get_app(transport: str = "streamable-http", enabled_apps: list[str] | None =
|
|||||||
logger.info(f"Cleaned up {count} expired login flow sessions")
|
logger.info(f"Cleaned up {count} expired login flow sessions")
|
||||||
# Also clean up expired AS proxy codes/sessions
|
# Also clean up expired AS proxy codes/sessions
|
||||||
_cleanup_expired_proxy_codes()
|
_cleanup_expired_proxy_codes()
|
||||||
|
# Clean up expired web provision sessions
|
||||||
|
_cleanup_expired_provision_sessions()
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning(f"Login flow cleanup error: {e}")
|
logger.warning(f"Login flow cleanup error: {e}")
|
||||||
await anyio.sleep(3600) # Every hour
|
await anyio.sleep(3600) # Every hour
|
||||||
|
|
||||||
@asynccontextmanager
|
@asynccontextmanager
|
||||||
async def _maybe_login_flow_cleanup():
|
async def _maybe_login_flow_cleanup(app: Starlette):
|
||||||
"""Start Login Flow cleanup task if enabled."""
|
"""Start Login Flow cleanup task and provision poll task group.
|
||||||
if settings.enable_login_flow:
|
|
||||||
async with anyio.create_task_group() as tg:
|
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)
|
tg.start_soon(_login_flow_cleanup_loop)
|
||||||
yield
|
# Share task group with provision routes for background polling
|
||||||
tg.cancel_scope.cancel()
|
found_app_mount = False
|
||||||
else:
|
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
|
yield
|
||||||
|
tg.cancel_scope.cancel()
|
||||||
|
|
||||||
@asynccontextmanager
|
@asynccontextmanager
|
||||||
async def _mcp_session_with_login_flow():
|
async def _mcp_session_with_login_flow(app: Starlette):
|
||||||
"""Start MCP session manager with optional Login Flow cleanup."""
|
"""Start MCP session manager with optional Login Flow cleanup."""
|
||||||
async with AsyncExitStack() as stack:
|
async with AsyncExitStack() as stack:
|
||||||
await stack.enter_async_context(mcp.session_manager.run())
|
await stack.enter_async_context(mcp.session_manager.run())
|
||||||
await stack.enter_async_context(_maybe_login_flow_cleanup())
|
await stack.enter_async_context(_maybe_login_flow_cleanup(app))
|
||||||
yield
|
yield
|
||||||
|
|
||||||
@asynccontextmanager
|
@asynccontextmanager
|
||||||
@@ -1659,7 +1684,7 @@ def get_app(transport: str = "streamable-http", enabled_apps: list[str] | None =
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Run MCP session manager and yield
|
# Run MCP session manager and yield
|
||||||
async with _mcp_session_with_login_flow():
|
async with _mcp_session_with_login_flow(app):
|
||||||
try:
|
try:
|
||||||
yield
|
yield
|
||||||
finally:
|
finally:
|
||||||
@@ -1803,9 +1828,9 @@ def get_app(transport: str = "streamable-http", enabled_apps: list[str] | None =
|
|||||||
break
|
break
|
||||||
|
|
||||||
# Determine authentication mode for background sync
|
# Determine authentication mode for background sync
|
||||||
# Multi-user BasicAuth: use app passwords via Astrolabe (NOT OAuth)
|
# Login Flow v2 and multi-user BasicAuth: use app passwords
|
||||||
# OAuth mode: use OAuth refresh tokens (NOT app passwords)
|
# OAuth mode (without Login Flow): use OAuth refresh tokens
|
||||||
use_basic_auth = not oauth_enabled
|
use_basic_auth = not oauth_enabled or settings.enable_login_flow
|
||||||
|
|
||||||
# Start background tasks using anyio TaskGroup
|
# Start background tasks using anyio TaskGroup
|
||||||
async with anyio.create_task_group() as tg:
|
async with anyio.create_task_group() as tg:
|
||||||
@@ -1841,7 +1866,7 @@ def get_app(transport: str = "streamable-http", enabled_apps: list[str] | None =
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Run MCP session manager and yield
|
# Run MCP session manager and yield
|
||||||
async with _mcp_session_with_login_flow():
|
async with _mcp_session_with_login_flow(app):
|
||||||
try:
|
try:
|
||||||
yield
|
yield
|
||||||
finally:
|
finally:
|
||||||
@@ -1860,7 +1885,7 @@ def get_app(transport: str = "streamable-http", enabled_apps: list[str] | None =
|
|||||||
"To enable, set NEXTCLOUD_OIDC_CLIENT_ID and NEXTCLOUD_OIDC_CLIENT_SECRET."
|
"To enable, set NEXTCLOUD_OIDC_CLIENT_ID and NEXTCLOUD_OIDC_CLIENT_SECRET."
|
||||||
)
|
)
|
||||||
# Just run MCP session manager without vector sync
|
# Just run MCP session manager without vector sync
|
||||||
async with _mcp_session_with_login_flow():
|
async with _mcp_session_with_login_flow(app):
|
||||||
yield
|
yield
|
||||||
|
|
||||||
else:
|
else:
|
||||||
@@ -1880,7 +1905,7 @@ def get_app(transport: str = "streamable-http", enabled_apps: list[str] | None =
|
|||||||
logger.warning(
|
logger.warning(
|
||||||
"Vector sync enabled but TOKEN_ENCRYPTION_KEY not set"
|
"Vector sync enabled but TOKEN_ENCRYPTION_KEY not set"
|
||||||
)
|
)
|
||||||
async with _mcp_session_with_login_flow():
|
async with _mcp_session_with_login_flow(app):
|
||||||
yield
|
yield
|
||||||
|
|
||||||
# Health check endpoints for Kubernetes probes
|
# Health check endpoints for Kubernetes probes
|
||||||
@@ -2331,6 +2356,15 @@ def get_app(transport: str = "streamable-http", enabled_apps: list[str] | None =
|
|||||||
),
|
),
|
||||||
]
|
]
|
||||||
|
|
||||||
|
# Login Flow v2 web provisioning (only when Login Flow is enabled)
|
||||||
|
if settings.enable_login_flow:
|
||||||
|
browser_routes += [
|
||||||
|
Route("/provision", provision_page, methods=["GET"]), # /app/provision
|
||||||
|
Route(
|
||||||
|
"/provision/status", provision_status, methods=["GET"]
|
||||||
|
), # /app/provision/status
|
||||||
|
]
|
||||||
|
|
||||||
# Add static files mount if directory exists
|
# Add static files mount if directory exists
|
||||||
static_dir = os.path.join(os.path.dirname(__file__), "auth", "static")
|
static_dir = os.path.join(os.path.dirname(__file__), "auth", "static")
|
||||||
if os.path.isdir(static_dir):
|
if os.path.isdir(static_dir):
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ The flow has two steps:
|
|||||||
|
|
||||||
import logging
|
import logging
|
||||||
import ssl
|
import ssl
|
||||||
|
from urllib.parse import urlparse, urlunparse
|
||||||
|
|
||||||
from pydantic import BaseModel, Field
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
@@ -18,6 +19,26 @@ from nextcloud_mcp_server.http import nextcloud_httpx_client
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def rewrite_url_origin(url: str, target_host: str) -> str:
|
||||||
|
"""Rewrite a URL's scheme+host+port to match target_host.
|
||||||
|
|
||||||
|
Preserves the path, params, query, and fragment from the original URL.
|
||||||
|
Useful for rewriting internal Docker hostnames to public-facing URLs.
|
||||||
|
"""
|
||||||
|
parsed_url = urlparse(url)
|
||||||
|
parsed_host = urlparse(target_host)
|
||||||
|
return urlunparse(
|
||||||
|
(
|
||||||
|
parsed_host.scheme,
|
||||||
|
parsed_host.netloc,
|
||||||
|
parsed_url.path,
|
||||||
|
parsed_url.params,
|
||||||
|
parsed_url.query,
|
||||||
|
parsed_url.fragment,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class LoginFlowInitResponse(BaseModel):
|
class LoginFlowInitResponse(BaseModel):
|
||||||
"""Response from initiating Login Flow v2."""
|
"""Response from initiating Login Flow v2."""
|
||||||
|
|
||||||
@@ -91,9 +112,16 @@ class LoginFlowV2Client:
|
|||||||
poll_data = data.get("poll", {})
|
poll_data = data.get("poll", {})
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
raw_poll_endpoint = poll_data["endpoint"]
|
||||||
|
# Nextcloud returns URLs using its internal hostname (e.g.
|
||||||
|
# http://localhost/login/v2/poll) which may be unreachable from
|
||||||
|
# this process. Rewrite the poll endpoint to use nextcloud_host
|
||||||
|
# so server-side polling works across Docker networks.
|
||||||
|
poll_endpoint = self._rewrite_to_nextcloud_host(raw_poll_endpoint)
|
||||||
|
|
||||||
result = LoginFlowInitResponse(
|
result = LoginFlowInitResponse(
|
||||||
login_url=data["login"],
|
login_url=data["login"],
|
||||||
poll_endpoint=poll_data["endpoint"],
|
poll_endpoint=poll_endpoint,
|
||||||
poll_token=poll_data["token"],
|
poll_token=poll_data["token"],
|
||||||
)
|
)
|
||||||
except KeyError as e:
|
except KeyError as e:
|
||||||
@@ -104,6 +132,19 @@ class LoginFlowV2Client:
|
|||||||
logger.info(f"Login Flow v2 initiated: login_url={result.login_url[:60]}...")
|
logger.info(f"Login Flow v2 initiated: login_url={result.login_url[:60]}...")
|
||||||
return result
|
return result
|
||||||
|
|
||||||
|
def _rewrite_to_nextcloud_host(self, url: str) -> str:
|
||||||
|
"""Rewrite a URL's origin to use self.nextcloud_host.
|
||||||
|
|
||||||
|
Nextcloud may return URLs with its internal hostname (e.g.
|
||||||
|
http://localhost) which differs from the configured NEXTCLOUD_HOST
|
||||||
|
(e.g. http://app:80). This replaces the scheme+host+port while
|
||||||
|
preserving the path and query.
|
||||||
|
"""
|
||||||
|
result = rewrite_url_origin(url, self.nextcloud_host)
|
||||||
|
if result != url:
|
||||||
|
logger.debug(f"Rewrote Login Flow v2 URL: {url} → {result}")
|
||||||
|
return result
|
||||||
|
|
||||||
async def poll(self, poll_endpoint: str, poll_token: str) -> LoginFlowPollResult:
|
async def poll(self, poll_endpoint: str, poll_token: str) -> LoginFlowPollResult:
|
||||||
"""Poll for Login Flow v2 completion by sending an HTTP POST to the Nextcloud instance.
|
"""Poll for Login Flow v2 completion by sending an HTTP POST to the Nextcloud instance.
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,346 @@
|
|||||||
|
"""Web-based Login Flow v2 provisioning routes.
|
||||||
|
|
||||||
|
Provides browser endpoints for provisioning Nextcloud app passwords via
|
||||||
|
Login Flow v2. Used by Astrolabe's "Enable Semantic Search" flow to
|
||||||
|
chain OAuth (bearer token) with Login Flow v2 (app password) in a single
|
||||||
|
user interaction.
|
||||||
|
|
||||||
|
Flow:
|
||||||
|
1. GET /app/provision?redirect_uri=... → Initiates LFv2, redirects to NC login
|
||||||
|
2. User clicks "Grant access" on Nextcloud's login page
|
||||||
|
3. MCP server background task polls and stores app password
|
||||||
|
4. GET /app/provision/status?id=... → Returns completion status (JSON)
|
||||||
|
5. User returns to Astrolabe settings (via redirect_uri or navigation)
|
||||||
|
"""
|
||||||
|
|
||||||
|
import html
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import secrets
|
||||||
|
import time
|
||||||
|
from urllib.parse import urlparse
|
||||||
|
|
||||||
|
import anyio
|
||||||
|
from starlette.requests import Request
|
||||||
|
from starlette.responses import HTMLResponse, JSONResponse, RedirectResponse
|
||||||
|
|
||||||
|
from nextcloud_mcp_server.api.management import validate_token_and_get_user
|
||||||
|
from nextcloud_mcp_server.auth.login_flow import LoginFlowV2Client, rewrite_url_origin
|
||||||
|
from nextcloud_mcp_server.auth.storage import get_shared_storage
|
||||||
|
from nextcloud_mcp_server.config import get_nextcloud_ssl_verify, get_settings
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# In-memory store for web provision sessions (short-lived, no persistence needed).
|
||||||
|
# Maps provision_id → session data.
|
||||||
|
# NOTE: This does not work with multi-process deployments (e.g. uvicorn --workers N).
|
||||||
|
# Login Flow v2 mode assumes a single worker process.
|
||||||
|
_provision_sessions: dict[str, dict] = {}
|
||||||
|
|
||||||
|
# Session TTL: 20 minutes (matches Nextcloud's Login Flow v2 timeout)
|
||||||
|
_SESSION_TTL = 1200
|
||||||
|
|
||||||
|
|
||||||
|
def _cleanup_expired_sessions() -> None:
|
||||||
|
"""Remove expired provision sessions."""
|
||||||
|
now = time.time()
|
||||||
|
expired = [k for k, v in _provision_sessions.items() if v["expires_at"] < now]
|
||||||
|
for k in expired:
|
||||||
|
del _provision_sessions[k]
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_redirect_uri(redirect_uri: str) -> bool:
|
||||||
|
"""Validate that redirect_uri is a reasonable URL (not javascript: etc)."""
|
||||||
|
try:
|
||||||
|
parsed = urlparse(redirect_uri)
|
||||||
|
return parsed.scheme in ("http", "https") and bool(parsed.netloc)
|
||||||
|
except Exception:
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
async def _poll_and_store(provision_id: str) -> None:
|
||||||
|
"""Background task: poll Login Flow v2 and store app password on completion."""
|
||||||
|
session = _provision_sessions.get(provision_id)
|
||||||
|
if not session:
|
||||||
|
return
|
||||||
|
|
||||||
|
settings = get_settings()
|
||||||
|
nextcloud_host = settings.nextcloud_host
|
||||||
|
if not nextcloud_host:
|
||||||
|
if provision_id in _provision_sessions:
|
||||||
|
session["status"] = "error"
|
||||||
|
return
|
||||||
|
|
||||||
|
flow_client = LoginFlowV2Client(
|
||||||
|
nextcloud_host=nextcloud_host,
|
||||||
|
verify_ssl=get_nextcloud_ssl_verify(),
|
||||||
|
)
|
||||||
|
|
||||||
|
poll_endpoint = session["poll_endpoint"]
|
||||||
|
poll_token = session["poll_token"]
|
||||||
|
user_id = session.get("user_id")
|
||||||
|
|
||||||
|
# Poll every 2 seconds for up to 20 minutes
|
||||||
|
max_attempts = 600
|
||||||
|
for _ in range(max_attempts):
|
||||||
|
if provision_id not in _provision_sessions:
|
||||||
|
return # Session was cleaned up
|
||||||
|
|
||||||
|
try:
|
||||||
|
result = await flow_client.poll(poll_endpoint, poll_token)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(
|
||||||
|
f"Login Flow v2 poll error for provision {provision_id}: {e}"
|
||||||
|
)
|
||||||
|
await anyio.sleep(2)
|
||||||
|
continue
|
||||||
|
|
||||||
|
if result.status == "completed":
|
||||||
|
# Store the app password
|
||||||
|
storage = await get_shared_storage()
|
||||||
|
effective_user_id = user_id or result.login_name or "unknown"
|
||||||
|
if not result.app_password:
|
||||||
|
# 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})"
|
||||||
|
)
|
||||||
|
return
|
||||||
|
await storage.store_app_password_with_scopes(
|
||||||
|
user_id=effective_user_id,
|
||||||
|
app_password=result.app_password,
|
||||||
|
scopes=None, # All scopes
|
||||||
|
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})"
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
if result.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})"
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
await anyio.sleep(2)
|
||||||
|
|
||||||
|
# Timed out
|
||||||
|
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})"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def provision_page(
|
||||||
|
request: Request,
|
||||||
|
) -> RedirectResponse | HTMLResponse | JSONResponse:
|
||||||
|
"""Initiate Login Flow v2 and redirect to Nextcloud's login page.
|
||||||
|
|
||||||
|
GET /app/provision?redirect_uri=...
|
||||||
|
|
||||||
|
Requires a valid Nextcloud OIDC bearer token (Authorization header).
|
||||||
|
The authenticated user identity is extracted from the token — the
|
||||||
|
``user_id`` query parameter is ignored if present.
|
||||||
|
|
||||||
|
Initiates Login Flow v2, starts background polling, and redirects the
|
||||||
|
browser to Nextcloud's login/grant page. After the user grants access,
|
||||||
|
the background task stores the app password. The user then navigates
|
||||||
|
back to the redirect_uri (Astrolabe settings).
|
||||||
|
"""
|
||||||
|
# Authenticate: require a valid Nextcloud OIDC bearer token
|
||||||
|
try:
|
||||||
|
user_id, _token_data = await validate_token_and_get_user(request)
|
||||||
|
except (ValueError, KeyError, AttributeError) as e:
|
||||||
|
logger.warning(f"Provision request rejected: {e}")
|
||||||
|
return JSONResponse({"error": "Authentication required"}, status_code=401)
|
||||||
|
|
||||||
|
_cleanup_expired_sessions()
|
||||||
|
|
||||||
|
redirect_uri = request.query_params.get("redirect_uri", "")
|
||||||
|
|
||||||
|
if not redirect_uri or not _validate_redirect_uri(redirect_uri):
|
||||||
|
return HTMLResponse(
|
||||||
|
content=_render_error("Missing or invalid redirect_uri parameter."),
|
||||||
|
status_code=400,
|
||||||
|
)
|
||||||
|
|
||||||
|
if urlparse(redirect_uri).scheme == "http":
|
||||||
|
logger.warning(f"Provision redirect_uri uses insecure HTTP: {redirect_uri}")
|
||||||
|
|
||||||
|
# Check if user already has an app password — skip straight to redirect
|
||||||
|
if user_id:
|
||||||
|
storage = await get_shared_storage()
|
||||||
|
existing = await storage.get_app_password_with_scopes(user_id)
|
||||||
|
if existing:
|
||||||
|
logger.info(f"User {user_id} already has app password, skipping provision")
|
||||||
|
return RedirectResponse(redirect_uri)
|
||||||
|
|
||||||
|
# Initiate Login Flow v2
|
||||||
|
settings = get_settings()
|
||||||
|
nextcloud_host = settings.nextcloud_host
|
||||||
|
if not nextcloud_host:
|
||||||
|
return HTMLResponse(
|
||||||
|
content=_render_error("Nextcloud host not configured on server."),
|
||||||
|
status_code=500,
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
flow_client = LoginFlowV2Client(
|
||||||
|
nextcloud_host=nextcloud_host,
|
||||||
|
verify_ssl=get_nextcloud_ssl_verify(),
|
||||||
|
)
|
||||||
|
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(
|
||||||
|
content=_render_error(
|
||||||
|
"Failed to start login flow. Please try again later."
|
||||||
|
),
|
||||||
|
status_code=502,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Create provision session
|
||||||
|
provision_id = secrets.token_urlsafe(32)
|
||||||
|
_provision_sessions[provision_id] = {
|
||||||
|
"status": "pending",
|
||||||
|
"login_url": init_response.login_url,
|
||||||
|
"poll_endpoint": init_response.poll_endpoint,
|
||||||
|
"poll_token": init_response.poll_token,
|
||||||
|
"redirect_uri": redirect_uri,
|
||||||
|
"user_id": user_id,
|
||||||
|
"created_at": time.time(),
|
||||||
|
"expires_at": time.time() + _SESSION_TTL,
|
||||||
|
}
|
||||||
|
|
||||||
|
# Start background polling task (uses task group from app lifespan)
|
||||||
|
poll_tg = getattr(request.app.state, "poll_task_group", None)
|
||||||
|
if poll_tg is None:
|
||||||
|
logger.error("No poll task group available; cannot start background polling")
|
||||||
|
_provision_sessions.pop(provision_id, None)
|
||||||
|
return HTMLResponse(
|
||||||
|
content=_render_error(
|
||||||
|
"Server configuration error: background polling unavailable."
|
||||||
|
),
|
||||||
|
status_code=500,
|
||||||
|
)
|
||||||
|
poll_tg.start_soon(_poll_and_store, provision_id)
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
f"Login Flow v2 web provision initiated (provision_id={provision_id}, "
|
||||||
|
f"user_id={user_id or 'unknown'}), redirecting to NC login"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Redirect to Nextcloud's Login Flow v2 login page.
|
||||||
|
# The login_url may use the internal Docker hostname (http://app/...).
|
||||||
|
# Replace with the public Nextcloud URL for the browser.
|
||||||
|
# Note: poll_endpoint is rewritten to NEXTCLOUD_HOST (server-side, in
|
||||||
|
# LoginFlowV2Client) while login_url is rewritten to the public issuer
|
||||||
|
# URL here because the browser needs a publicly-reachable address.
|
||||||
|
login_url = init_response.login_url
|
||||||
|
public_issuer = os.getenv("NEXTCLOUD_PUBLIC_ISSUER_URL", "")
|
||||||
|
if public_issuer and nextcloud_host:
|
||||||
|
login_url = rewrite_url_origin(login_url, public_issuer.rstrip("/"))
|
||||||
|
|
||||||
|
return RedirectResponse(login_url)
|
||||||
|
|
||||||
|
|
||||||
|
async def provision_status(request: Request) -> JSONResponse:
|
||||||
|
"""Check provision session status.
|
||||||
|
|
||||||
|
GET /app/provision/status?id=...
|
||||||
|
|
||||||
|
Requires a valid Nextcloud OIDC bearer token (Authorization header).
|
||||||
|
|
||||||
|
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)
|
||||||
|
"""
|
||||||
|
# Authenticate: require a valid Nextcloud OIDC bearer token
|
||||||
|
try:
|
||||||
|
_user_id, _token_data = await validate_token_and_get_user(request)
|
||||||
|
except (ValueError, KeyError, AttributeError) as e:
|
||||||
|
logger.warning(f"Provision status request rejected: {e}")
|
||||||
|
return JSONResponse({"error": "Authentication required"}, status_code=401)
|
||||||
|
|
||||||
|
provision_id = request.query_params.get("id", "")
|
||||||
|
|
||||||
|
session = _provision_sessions.get(provision_id)
|
||||||
|
if not session:
|
||||||
|
return JSONResponse(
|
||||||
|
{
|
||||||
|
"status": "not_found",
|
||||||
|
"message": "Provision session not found or expired",
|
||||||
|
},
|
||||||
|
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")
|
||||||
|
# Clean up completed session after status is read
|
||||||
|
_provision_sessions.pop(provision_id, None)
|
||||||
|
|
||||||
|
return JSONResponse(response)
|
||||||
|
|
||||||
|
|
||||||
|
# ── HTML rendering helpers ────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
def _render_error(message: str) -> str:
|
||||||
|
"""Render an error page."""
|
||||||
|
return f"""<!DOCTYPE html>
|
||||||
|
<html lang="en">
|
||||||
|
<head>
|
||||||
|
<meta charset="UTF-8">
|
||||||
|
<meta name="viewport" content="width=device-width, initial-scale=1.0">
|
||||||
|
<title>Error - Astrolabe</title>
|
||||||
|
<style>
|
||||||
|
body {{
|
||||||
|
font-family: -apple-system, BlinkMacSystemFont, 'Segoe UI', Roboto, sans-serif;
|
||||||
|
background: #f5f5f5;
|
||||||
|
display: flex;
|
||||||
|
justify-content: center;
|
||||||
|
align-items: center;
|
||||||
|
min-height: 100vh;
|
||||||
|
}}
|
||||||
|
.card {{
|
||||||
|
background: #fff;
|
||||||
|
border-radius: 12px;
|
||||||
|
box-shadow: 0 2px 8px rgba(0,0,0,0.1);
|
||||||
|
padding: 2.5rem;
|
||||||
|
max-width: 480px;
|
||||||
|
text-align: center;
|
||||||
|
}}
|
||||||
|
.error {{ color: #c62828; }}
|
||||||
|
</style>
|
||||||
|
</head>
|
||||||
|
<body>
|
||||||
|
<div class="card">
|
||||||
|
<h1 class="error">Provisioning Error</h1>
|
||||||
|
<p>{html.escape(message)}</p>
|
||||||
|
</div>
|
||||||
|
</body>
|
||||||
|
</html>"""
|
||||||
@@ -489,7 +489,7 @@ async def user_manager_task(
|
|||||||
try:
|
try:
|
||||||
# Get current provisioned users based on mode
|
# Get current provisioned users based on mode
|
||||||
if use_basic_auth:
|
if use_basic_auth:
|
||||||
# BasicAuth mode: query app_passwords table
|
# BasicAuth / Login Flow v2 mode: query app_passwords table
|
||||||
provisioned_users = set(
|
provisioned_users = set(
|
||||||
await refresh_token_storage.get_all_app_password_user_ids()
|
await refresh_token_storage.get_all_app_password_user_ids()
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ from nextcloud_mcp_server.auth.login_flow import (
|
|||||||
LoginFlowInitResponse,
|
LoginFlowInitResponse,
|
||||||
LoginFlowPollResult,
|
LoginFlowPollResult,
|
||||||
LoginFlowV2Client,
|
LoginFlowV2Client,
|
||||||
|
rewrite_url_origin,
|
||||||
)
|
)
|
||||||
|
|
||||||
pytestmark = pytest.mark.unit
|
pytestmark = pytest.mark.unit
|
||||||
@@ -208,3 +209,36 @@ async def test_login_flow_poll_result_model():
|
|||||||
assert pending.status == "pending"
|
assert pending.status == "pending"
|
||||||
assert pending.server is None
|
assert pending.server is None
|
||||||
assert pending.app_password is None
|
assert pending.app_password is None
|
||||||
|
|
||||||
|
|
||||||
|
# ── rewrite_url_origin tests ─────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
async def test_rewrite_url_origin_basic():
|
||||||
|
"""Test basic origin rewriting."""
|
||||||
|
result = rewrite_url_origin(
|
||||||
|
"http://localhost/login/v2/poll", "https://cloud.example.com"
|
||||||
|
)
|
||||||
|
assert result == "https://cloud.example.com/login/v2/poll"
|
||||||
|
|
||||||
|
|
||||||
|
async def test_rewrite_url_origin_preserves_port():
|
||||||
|
"""Test that port in target_host is preserved."""
|
||||||
|
result = rewrite_url_origin("http://localhost/path", "http://app:8080")
|
||||||
|
assert result == "http://app:8080/path"
|
||||||
|
|
||||||
|
|
||||||
|
async def test_rewrite_url_origin_preserves_query():
|
||||||
|
"""Test that query string and fragment are preserved."""
|
||||||
|
result = rewrite_url_origin(
|
||||||
|
"http://internal/path?token=abc&foo=bar#section",
|
||||||
|
"https://public.example.com",
|
||||||
|
)
|
||||||
|
assert result == "https://public.example.com/path?token=abc&foo=bar#section"
|
||||||
|
|
||||||
|
|
||||||
|
async def test_rewrite_url_origin_noop_when_same():
|
||||||
|
"""Test that rewriting to the same origin is a no-op."""
|
||||||
|
url = "https://cloud.example.com/login/v2/poll"
|
||||||
|
result = rewrite_url_origin(url, "https://cloud.example.com")
|
||||||
|
assert result == url
|
||||||
|
|||||||
@@ -0,0 +1,409 @@
|
|||||||
|
"""Unit tests for web-based Login Flow v2 provisioning routes.
|
||||||
|
|
||||||
|
Tests validation, HTML escaping, URL rewriting, route handlers, and
|
||||||
|
background polling logic.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import time
|
||||||
|
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,
|
||||||
|
provision_page,
|
||||||
|
provision_status,
|
||||||
|
)
|
||||||
|
|
||||||
|
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 ─────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
async def test_validate_redirect_uri_accepts_https():
|
||||||
|
"""Valid HTTPS URL is accepted."""
|
||||||
|
assert _validate_redirect_uri("https://app.example.com/callback") is True
|
||||||
|
|
||||||
|
|
||||||
|
async def test_validate_redirect_uri_accepts_http_localhost():
|
||||||
|
"""Valid HTTP localhost URL is accepted."""
|
||||||
|
assert _validate_redirect_uri("http://localhost:3000/callback") is True
|
||||||
|
|
||||||
|
|
||||||
|
async def test_validate_redirect_uri_rejects_javascript():
|
||||||
|
"""javascript: URIs are rejected."""
|
||||||
|
assert _validate_redirect_uri("javascript:alert(1)") is False
|
||||||
|
|
||||||
|
|
||||||
|
async def test_validate_redirect_uri_rejects_relative_url():
|
||||||
|
"""Relative URLs are rejected (no scheme/netloc)."""
|
||||||
|
assert _validate_redirect_uri("/relative/path") is False
|
||||||
|
|
||||||
|
|
||||||
|
async def test_validate_redirect_uri_rejects_bare_hostname():
|
||||||
|
"""Bare hostnames without scheme are rejected."""
|
||||||
|
assert _validate_redirect_uri("example.com") is False
|
||||||
|
|
||||||
|
|
||||||
|
async def test_validate_redirect_uri_rejects_data_uri():
|
||||||
|
"""data: URIs are rejected."""
|
||||||
|
assert _validate_redirect_uri("data:text/html,<h1>hi</h1>") is False
|
||||||
|
|
||||||
|
|
||||||
|
async def test_validate_redirect_uri_rejects_empty():
|
||||||
|
"""Empty string is rejected."""
|
||||||
|
assert _validate_redirect_uri("") is False
|
||||||
|
|
||||||
|
|
||||||
|
# ── _render_error tests ──────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
async def test_render_error_escapes_html():
|
||||||
|
"""XSS regression: HTML in error messages must be escaped."""
|
||||||
|
html_output = _render_error("<script>alert('xss')</script>")
|
||||||
|
assert "<script>" not in html_output
|
||||||
|
assert "<script>" in html_output
|
||||||
|
|
||||||
|
|
||||||
|
async def test_render_error_escapes_angle_brackets():
|
||||||
|
"""Angle brackets in exception messages are escaped."""
|
||||||
|
html_output = _render_error("Unexpected response from <internal-host>")
|
||||||
|
assert "<internal-host>" not in html_output
|
||||||
|
assert "<internal-host>" in html_output
|
||||||
|
|
||||||
|
|
||||||
|
async def test_render_error_preserves_plain_text():
|
||||||
|
"""Plain text messages render correctly."""
|
||||||
|
html_output = _render_error("Something went wrong.")
|
||||||
|
assert "Something went wrong." in html_output
|
||||||
|
assert "Provisioning Error" in html_output
|
||||||
|
|
||||||
|
|
||||||
|
# ── Auth helper ──────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
_MOCK_TOKEN_PATCH = patch(
|
||||||
|
"nextcloud_mcp_server.auth.provision_routes.validate_token_and_get_user",
|
||||||
|
new_callable=AsyncMock,
|
||||||
|
return_value=("alice", {"sub": "alice", "client_id": "astrolabe", "scopes": []}),
|
||||||
|
)
|
||||||
|
"""Patch that makes validate_token_and_get_user succeed as user 'alice'."""
|
||||||
|
|
||||||
|
|
||||||
|
# ── provision_status tests ───────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
def _make_request(query_params: dict) -> MagicMock:
|
||||||
|
"""Create a mock Starlette Request with query_params."""
|
||||||
|
request = MagicMock()
|
||||||
|
request.query_params = query_params
|
||||||
|
return request
|
||||||
|
|
||||||
|
|
||||||
|
async def test_provision_status_rejects_missing_token():
|
||||||
|
"""Missing bearer token returns 401."""
|
||||||
|
request = _make_request({"id": "some-id"})
|
||||||
|
with patch(
|
||||||
|
"nextcloud_mcp_server.auth.provision_routes.validate_token_and_get_user",
|
||||||
|
new_callable=AsyncMock,
|
||||||
|
side_effect=ValueError("Missing Authorization header"),
|
||||||
|
):
|
||||||
|
response = await provision_status(request)
|
||||||
|
assert response.status_code == 401
|
||||||
|
|
||||||
|
|
||||||
|
async def test_provision_status_not_found():
|
||||||
|
"""Unknown provision ID returns 404."""
|
||||||
|
request = _make_request({"id": "nonexistent-id"})
|
||||||
|
with _MOCK_TOKEN_PATCH:
|
||||||
|
response = await provision_status(request)
|
||||||
|
assert response.status_code == 404
|
||||||
|
assert response.body is not None
|
||||||
|
|
||||||
|
|
||||||
|
async def test_provision_status_pending():
|
||||||
|
"""Pending session returns status=pending."""
|
||||||
|
provision_id = "test-pending-id"
|
||||||
|
_provision_sessions[provision_id] = {
|
||||||
|
"status": "pending",
|
||||||
|
"expires_at": time.time() + 600,
|
||||||
|
}
|
||||||
|
request = _make_request({"id": provision_id})
|
||||||
|
with _MOCK_TOKEN_PATCH:
|
||||||
|
response = await provision_status(request)
|
||||||
|
assert response.status_code == 200
|
||||||
|
|
||||||
|
|
||||||
|
async def test_provision_status_completed_cleans_up():
|
||||||
|
"""Completed session returns username and removes session."""
|
||||||
|
provision_id = "test-completed-id"
|
||||||
|
_provision_sessions[provision_id] = {
|
||||||
|
"status": "completed",
|
||||||
|
"username": "alice",
|
||||||
|
"expires_at": time.time() + 600,
|
||||||
|
}
|
||||||
|
request = _make_request({"id": provision_id})
|
||||||
|
with _MOCK_TOKEN_PATCH:
|
||||||
|
response = await provision_status(request)
|
||||||
|
assert response.status_code == 200
|
||||||
|
# Session should be cleaned up after status read
|
||||||
|
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})
|
||||||
|
with _MOCK_TOKEN_PATCH:
|
||||||
|
response = await provision_status(request)
|
||||||
|
assert response.status_code == 404
|
||||||
|
assert provision_id not in _provision_sessions
|
||||||
|
|
||||||
|
|
||||||
|
# ── provision_page tests ─────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
async def test_provision_page_rejects_missing_token():
|
||||||
|
"""Missing bearer token returns 401."""
|
||||||
|
request = _make_request({"redirect_uri": "https://example.com/callback"})
|
||||||
|
with patch(
|
||||||
|
"nextcloud_mcp_server.auth.provision_routes.validate_token_and_get_user",
|
||||||
|
new_callable=AsyncMock,
|
||||||
|
side_effect=ValueError("Missing Authorization header"),
|
||||||
|
):
|
||||||
|
response = await provision_page(request)
|
||||||
|
assert response.status_code == 401
|
||||||
|
|
||||||
|
|
||||||
|
async def test_provision_page_rejects_invalid_token():
|
||||||
|
"""Invalid bearer token returns 401."""
|
||||||
|
request = _make_request({"redirect_uri": "https://example.com/callback"})
|
||||||
|
with patch(
|
||||||
|
"nextcloud_mcp_server.auth.provision_routes.validate_token_and_get_user",
|
||||||
|
new_callable=AsyncMock,
|
||||||
|
side_effect=ValueError("Token validation failed"),
|
||||||
|
):
|
||||||
|
response = await provision_page(request)
|
||||||
|
assert response.status_code == 401
|
||||||
|
|
||||||
|
|
||||||
|
async def test_provision_page_missing_redirect_uri():
|
||||||
|
"""Missing redirect_uri returns 400."""
|
||||||
|
request = _make_request({})
|
||||||
|
with _MOCK_TOKEN_PATCH:
|
||||||
|
response = await provision_page(request)
|
||||||
|
assert response.status_code == 400
|
||||||
|
|
||||||
|
|
||||||
|
async def test_provision_page_invalid_redirect_uri():
|
||||||
|
"""Invalid redirect_uri (javascript:) returns 400."""
|
||||||
|
request = _make_request({"redirect_uri": "javascript:alert(1)"})
|
||||||
|
with _MOCK_TOKEN_PATCH:
|
||||||
|
response = await provision_page(request)
|
||||||
|
assert response.status_code == 400
|
||||||
|
|
||||||
|
|
||||||
|
async def test_provision_page_skips_if_already_provisioned():
|
||||||
|
"""If user already has an app password, redirect immediately."""
|
||||||
|
request = _make_request(
|
||||||
|
{
|
||||||
|
"redirect_uri": "https://app.example.com/settings",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
mock_storage = AsyncMock()
|
||||||
|
mock_storage.get_app_password_with_scopes.return_value = {
|
||||||
|
"app_password": "existing-password",
|
||||||
|
}
|
||||||
|
|
||||||
|
with (
|
||||||
|
_MOCK_TOKEN_PATCH,
|
||||||
|
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