diff --git a/frontends/conductor.py b/frontends/conductor.py index 9202ce7dc..0d9193fa3 100644 --- a/frontends/conductor.py +++ b/frontends/conductor.py @@ -2,10 +2,10 @@ from dataclasses import dataclass, field from typing import Dict, Any, Optional, List from contextlib import asynccontextmanager, suppress +from urllib.parse import urlsplit from fastapi import FastAPI, WebSocket, WebSocketDisconnect, Request, HTTPException from fastapi.responses import FileResponse, PlainTextResponse, JSONResponse -from fastapi.middleware.cors import CORSMiddleware from pydantic import BaseModel def _resolve_ga_root() -> str: @@ -118,8 +118,81 @@ async def lifespan(app: FastAPI): yield +_LOOPBACK_HOSTS = {"127.0.0.1", "localhost", "::1"} + + +def _origin_hostname(origin: str) -> str: + if origin != origin.strip(): return "" + try: + parsed = urlsplit(origin) + parsed.port + except ValueError: + return "" + if (parsed.scheme.lower() not in ("http", "https") or not parsed.hostname or + parsed.username is not None or parsed.password is not None or + parsed.path not in ("", "/") or parsed.query or parsed.fragment): + return "" + return parsed.hostname.lower() + + +def _request_hostname(host: str) -> str: + try: + parsed = urlsplit("//" + host) + parsed.port + except ValueError: + return "" + if parsed.username is not None or parsed.password is not None: + return "" + return (parsed.hostname or "").lower() + + +def _origin_matches_host(origin: str, host: str) -> bool: + origin_host = _origin_hostname(origin) + request_host = _request_hostname(host) + if not origin_host or not request_host: + return False + if args.host.lower() in _LOOPBACK_HOSTS: + return origin_host in _LOOPBACK_HOSTS and request_host in _LOOPBACK_HOSTS + return origin_host == request_host + + +def _cors_headers(origin: str) -> dict: + return { + "Access-Control-Allow-Origin": origin, + "Access-Control-Allow-Methods": "GET,POST,PUT,PATCH,DELETE,OPTIONS", + "Access-Control-Allow-Headers": "Content-Type,Authorization", + "Vary": "Origin", + } + + +class BrowserOriginGuard: + def __init__(self, app): self.app = app + async def __call__(self, scope, receive, send): + if scope.get("type") not in ("http", "websocket"): + return await self.app(scope, receive, send) + headers = {k.decode("latin-1").lower(): v.decode("latin-1") for k, v in scope.get("headers", [])} + origin, host = headers.get("origin", ""), headers.get("host", "") + if origin and not _origin_matches_host(origin, host): + if scope["type"] == "websocket": + return await send({"type": "websocket.close", "code": 1008}) + return await PlainTextResponse("Forbidden", status_code=403)(scope, receive, send) + if scope["type"] == "websocket": + return await self.app(scope, receive, send) + cors = _cors_headers(origin) if origin else {} + if scope.get("method") == "OPTIONS" and origin: + return await PlainTextResponse("", status_code=204, headers=cors)(scope, receive, send) + async def send_with_cors(message): + if cors and message.get("type") == "http.response.start": + message = dict(message) + message["headers"] = list(message.get("headers", [])) + [ + (k.encode("latin-1"), v.encode("latin-1")) for k, v in cors.items() + ] + await send(message) + await self.app(scope, receive, send_with_cors) + + app = FastAPI(title="Conductor", lifespan=lifespan) -app.add_middleware(CORSMiddleware, allow_origins=["*"], allow_methods=["*"], allow_headers=["*"]) +app.add_middleware(BrowserOriginGuard) class RemoteAuth: def __init__(self, app): self.app = app diff --git a/tests/test_conductor_origin.py b/tests/test_conductor_origin.py new file mode 100644 index 000000000..50c5bd94a --- /dev/null +++ b/tests/test_conductor_origin.py @@ -0,0 +1,119 @@ +import importlib.util +import sys +import unittest +from pathlib import Path + +from fastapi import FastAPI, WebSocket +from fastapi.testclient import TestClient +from starlette.websockets import WebSocketDisconnect + + +ROOT = Path(__file__).resolve().parents[1] +FRONTENDS = ROOT / "frontends" +if str(ROOT) not in sys.path: + sys.path.insert(0, str(ROOT)) + +_old_argv = sys.argv[:] +try: + sys.argv = [sys.argv[0], "--no-browser"] + spec = importlib.util.spec_from_file_location("conductor", FRONTENDS / "conductor.py") + assert spec and spec.loader + conductor = importlib.util.module_from_spec(spec) + sys.modules[spec.name] = conductor + spec.loader.exec_module(conductor) +finally: + sys.argv = _old_argv + + +def make_app(): + app = FastAPI() + app.add_middleware(conductor.BrowserOriginGuard) + + @app.post("/probe") + async def probe(): + return {"ok": True} + + @app.websocket("/ws") + async def websocket(ws: WebSocket): + await ws.accept() + await ws.send_text("ok") + await ws.close() + + return app + + +class ConductorBrowserOriginTests(unittest.TestCase): + def setUp(self): + self.client = TestClient(make_app(), base_url="http://127.0.0.1:8900") + + def test_cross_origin_preflight_is_rejected(self): + r = self.client.options( + "/probe", + headers={ + "Origin": "https://evil.example", + "Access-Control-Request-Method": "POST", + "Access-Control-Request-Headers": "content-type", + }, + ) + self.assertEqual(r.status_code, 403) + + def test_cross_origin_post_is_rejected(self): + r = self.client.post("/probe", headers={"Origin": "https://evil.example"}) + self.assertEqual(r.status_code, 403) + + def test_dns_rebinding_origin_is_rejected_on_loopback(self): + r = self.client.post( + "/probe", + headers={"Host": "evil.example:8900", "Origin": "https://evil.example"}, + ) + self.assertEqual(r.status_code, 403) + + def test_malformed_serialized_origins_are_rejected(self): + origins = ( + "http://127.0.0.1:14168/path", + "http://127.0.0.1:14168?x=1", + "http://user@127.0.0.1:14168", + "http://127.0.0.1:notaport", + "http://127.0.0.1:14168#fragment", + ) + for origin in origins: + with self.subTest(origin=origin): + r = self.client.post( + "/probe", + headers={"Host": "127.0.0.1:8900", "Origin": origin}, + ) + self.assertEqual(r.status_code, 403) + + def test_desktop_origin_same_hostname_different_port_is_allowed(self): + origin = "http://127.0.0.1:14168" + r = self.client.post("/probe", headers={"Origin": origin}) + self.assertEqual(r.status_code, 200) + self.assertEqual(r.headers.get("Access-Control-Allow-Origin"), origin) + + def test_null_origin_is_rejected(self): + r = self.client.post("/probe", headers={"Origin": "null"}) + self.assertEqual(r.status_code, 403) + + def test_non_browser_request_without_origin_is_allowed(self): + r = self.client.post("/probe") + self.assertEqual(r.status_code, 200) + self.assertNotIn("Access-Control-Allow-Origin", r.headers) + + def test_cross_origin_websocket_is_rejected(self): + with self.assertRaises(WebSocketDisconnect): + with self.client.websocket_connect( + "/ws", + headers={"Host": "127.0.0.1:8900", "Origin": "https://evil.example"}, + ): + pass + + def test_desktop_origin_websocket_is_allowed(self): + with self.client.websocket_connect( + "/ws", + headers={"Host": "127.0.0.1:8900", "Origin": "http://127.0.0.1:14168"}, + ) as ws: + self.assertEqual(ws.receive_text(), "ok") + + +if __name__ == "__main__": + unittest.main()