diff --git a/nextcloud_mcp_server/app.py b/nextcloud_mcp_server/app.py index 55028f23..f755e137 100644 --- a/nextcloud_mcp_server/app.py +++ b/nextcloud_mcp_server/app.py @@ -1412,22 +1412,26 @@ def get_app(transport: str = "streamable-http", enabled_apps: list[str] | None = 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.""" + 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 + 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 + break 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 @@ -1663,7 +1667,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: @@ -1845,7 +1849,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: @@ -1864,7 +1868,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: @@ -1884,7 +1888,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 diff --git a/nextcloud_mcp_server/auth/login_flow.py b/nextcloud_mcp_server/auth/login_flow.py index a6120611..0bfdd10a 100644 --- a/nextcloud_mcp_server/auth/login_flow.py +++ b/nextcloud_mcp_server/auth/login_flow.py @@ -19,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.""" @@ -120,13 +140,7 @@ class LoginFlowV2Client: (e.g. http://app:80). This replaces the scheme+host+port while preserving the path and query. """ - parsed_url = urlparse(url) - parsed_host = urlparse(self.nextcloud_host) - rewritten = parsed_url._replace( - scheme=parsed_host.scheme, - netloc=parsed_host.netloc, - ) - result = urlunparse(rewritten) + result = rewrite_url_origin(url, self.nextcloud_host) if result != url: logger.debug(f"Rewrote Login Flow v2 URL: {url} → {result}") return result diff --git a/nextcloud_mcp_server/auth/provision_routes.py b/nextcloud_mcp_server/auth/provision_routes.py index 16d9563e..179b6901 100644 --- a/nextcloud_mcp_server/auth/provision_routes.py +++ b/nextcloud_mcp_server/auth/provision_routes.py @@ -13,17 +13,18 @@ Flow: 5. User returns to Astrolabe settings (via redirect_uri or navigation) """ -import asyncio +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.auth.login_flow import LoginFlowV2Client +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 @@ -87,7 +88,7 @@ async def _poll_and_store(provision_id: str) -> None: logger.warning( f"Login Flow v2 poll error for provision {provision_id}: {e}" ) - await asyncio.sleep(2) + await anyio.sleep(2) continue if result.status == "completed": @@ -121,7 +122,7 @@ async def _poll_and_store(provision_id: str) -> None: ) return - await asyncio.sleep(2) + await anyio.sleep(2) # Timed out session["status"] = "expired" @@ -151,6 +152,9 @@ async def provision_page(request: Request) -> RedirectResponse | HTMLResponse: 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() @@ -196,8 +200,18 @@ async def provision_page(request: Request) -> RedirectResponse | HTMLResponse: "expires_at": time.time() + _SESSION_TTL, } - # Start background polling task - asyncio.create_task(_poll_and_store(provision_id)) + # 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}, " @@ -210,12 +224,7 @@ async def provision_page(request: Request) -> RedirectResponse | HTMLResponse: login_url = init_response.login_url public_issuer = os.getenv("NEXTCLOUD_PUBLIC_ISSUER_URL", "") if public_issuer and nextcloud_host: - # Extract just scheme+host from NEXTCLOUD_HOST for matching - # (NC may omit default ports, e.g. http://app:80 → http://app) - parsed = urlparse(nextcloud_host) - internal_origin = f"{parsed.scheme}://{parsed.hostname}" - if internal_origin in login_url: - login_url = login_url.replace(internal_origin, public_issuer.rstrip("/")) + login_url = rewrite_url_origin(login_url, public_issuer.rstrip("/")) return RedirectResponse(login_url) @@ -282,7 +291,7 @@ def _render_error(message: str) -> str:

Provisioning Error

-

{message}

+

{html.escape(message)}

""" diff --git a/tests/unit/test_login_flow.py b/tests/unit/test_login_flow.py index 6c7c51a0..a5dfc9bc 100644 --- a/tests/unit/test_login_flow.py +++ b/tests/unit/test_login_flow.py @@ -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 diff --git a/tests/unit/test_provision_routes.py b/tests/unit/test_provision_routes.py new file mode 100644 index 00000000..5cd67b16 --- /dev/null +++ b/tests/unit/test_provision_routes.py @@ -0,0 +1,170 @@ +"""Unit tests for web-based Login Flow v2 provisioning routes. + +Tests validation, HTML escaping, URL rewriting, and route handlers. +""" + +import time +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from nextcloud_mcp_server.auth.provision_routes import ( + _provision_sessions, + _render_error, + _validate_redirect_uri, + provision_page, + provision_status, +) + +pytestmark = pytest.mark.unit + + +# ── _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,

hi

") 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("") + assert "