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/
|
||||
|
||||
# Install dependencies
|
||||
# 1. git (required for caldav dependency from git)
|
||||
# 2. sqlite for development with token db
|
||||
# 1. curl (required for container healthcheck probes)
|
||||
# 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 \
|
||||
curl \
|
||||
git \
|
||||
tesseract-ocr \
|
||||
sqlite3 && apt clean
|
||||
|
||||
+50
-16
@@ -73,6 +73,10 @@ from nextcloud_mcp_server.auth.oauth_routes import (
|
||||
oauth_register_proxy,
|
||||
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.storage import RefreshTokenStorage, get_shared_storage
|
||||
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
|
||||
_cleanup_expired_proxy_codes,
|
||||
)
|
||||
from nextcloud_mcp_server.auth.provision_routes import ( # noqa: PLC0415
|
||||
_cleanup_expired_sessions as _cleanup_expired_provision_sessions,
|
||||
)
|
||||
|
||||
while True:
|
||||
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")
|
||||
# Also clean up expired AS proxy codes/sessions
|
||||
_cleanup_expired_proxy_codes()
|
||||
# Clean up expired web provision sessions
|
||||
_cleanup_expired_provision_sessions()
|
||||
except Exception as e:
|
||||
logger.warning(f"Login flow cleanup error: {e}")
|
||||
await anyio.sleep(3600) # Every hour
|
||||
|
||||
@asynccontextmanager
|
||||
async def _maybe_login_flow_cleanup():
|
||||
"""Start Login Flow cleanup task if enabled."""
|
||||
if settings.enable_login_flow:
|
||||
async with anyio.create_task_group() as tg:
|
||||
async def _maybe_login_flow_cleanup(app: Starlette):
|
||||
"""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)
|
||||
yield
|
||||
tg.cancel_scope.cancel()
|
||||
else:
|
||||
# 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()
|
||||
|
||||
@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."""
|
||||
async with AsyncExitStack() as stack:
|
||||
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
|
||||
|
||||
@asynccontextmanager
|
||||
@@ -1659,7 +1684,7 @@ def get_app(transport: str = "streamable-http", enabled_apps: list[str] | None =
|
||||
)
|
||||
|
||||
# Run MCP session manager and yield
|
||||
async with _mcp_session_with_login_flow():
|
||||
async with _mcp_session_with_login_flow(app):
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
@@ -1803,9 +1828,9 @@ def get_app(transport: str = "streamable-http", enabled_apps: list[str] | None =
|
||||
break
|
||||
|
||||
# Determine authentication mode for background sync
|
||||
# Multi-user BasicAuth: use app passwords via Astrolabe (NOT OAuth)
|
||||
# OAuth mode: use OAuth refresh tokens (NOT app passwords)
|
||||
use_basic_auth = not oauth_enabled
|
||||
# Login Flow v2 and multi-user BasicAuth: use app passwords
|
||||
# OAuth mode (without Login Flow): use OAuth refresh tokens
|
||||
use_basic_auth = not oauth_enabled or settings.enable_login_flow
|
||||
|
||||
# Start background tasks using anyio TaskGroup
|
||||
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
|
||||
async with _mcp_session_with_login_flow():
|
||||
async with _mcp_session_with_login_flow(app):
|
||||
try:
|
||||
yield
|
||||
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."
|
||||
)
|
||||
# Just run MCP session manager without vector sync
|
||||
async with _mcp_session_with_login_flow():
|
||||
async with _mcp_session_with_login_flow(app):
|
||||
yield
|
||||
|
||||
else:
|
||||
@@ -1880,7 +1905,7 @@ def get_app(transport: str = "streamable-http", enabled_apps: list[str] | None =
|
||||
logger.warning(
|
||||
"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
|
||||
|
||||
# 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
|
||||
static_dir = os.path.join(os.path.dirname(__file__), "auth", "static")
|
||||
if os.path.isdir(static_dir):
|
||||
|
||||
@@ -10,6 +10,7 @@ The flow has two steps:
|
||||
|
||||
import logging
|
||||
import ssl
|
||||
from urllib.parse import urlparse, urlunparse
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
@@ -18,6 +19,26 @@ from nextcloud_mcp_server.http import nextcloud_httpx_client
|
||||
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):
|
||||
"""Response from initiating Login Flow v2."""
|
||||
|
||||
@@ -91,9 +112,16 @@ class LoginFlowV2Client:
|
||||
poll_data = data.get("poll", {})
|
||||
|
||||
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(
|
||||
login_url=data["login"],
|
||||
poll_endpoint=poll_data["endpoint"],
|
||||
poll_endpoint=poll_endpoint,
|
||||
poll_token=poll_data["token"],
|
||||
)
|
||||
except KeyError as e:
|
||||
@@ -104,6 +132,19 @@ class LoginFlowV2Client:
|
||||
logger.info(f"Login Flow v2 initiated: login_url={result.login_url[:60]}...")
|
||||
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:
|
||||
"""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:
|
||||
# Get current provisioned users based on mode
|
||||
if use_basic_auth:
|
||||
# BasicAuth mode: query app_passwords table
|
||||
# BasicAuth / Login Flow v2 mode: query app_passwords table
|
||||
provisioned_users = set(
|
||||
await refresh_token_storage.get_all_app_password_user_ids()
|
||||
)
|
||||
|
||||
@@ -14,6 +14,7 @@ from nextcloud_mcp_server.auth.login_flow import (
|
||||
LoginFlowInitResponse,
|
||||
LoginFlowPollResult,
|
||||
LoginFlowV2Client,
|
||||
rewrite_url_origin,
|
||||
)
|
||||
|
||||
pytestmark = pytest.mark.unit
|
||||
@@ -208,3 +209,36 @@ async def test_login_flow_poll_result_model():
|
||||
assert pending.status == "pending"
|
||||
assert pending.server 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