Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
64 changes: 58 additions & 6 deletions frontends/desktop_bridge.py
Original file line number Diff line number Diff line change
Expand Up @@ -1399,26 +1399,78 @@ async def ws_handler(request):
# Transport layer: HTTP command/data API
# ---------------------------------------------------------------------------

def cors_headers():
_LOOPBACK_HOSTS = {"127.0.0.1", "localhost", "::1"}


def _origin_parts(origin: str):
if origin != origin.strip(): return None
from urllib.parse import urlsplit
try:
parsed = urlsplit(origin)
parsed.port
except ValueError:
return None
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 or parsed.query or parsed.fragment):
return None
return parsed


def _request_host_parts(host: str):
from urllib.parse import urlsplit
try:
parsed = urlsplit("//" + host)
parsed.port
except ValueError:
return None
if parsed.username is not None or parsed.password is not None or parsed.query or parsed.fragment:
return None
return parsed if parsed.hostname else None


def _origin_matches_host(request, origin: str) -> bool:
if not origin:
return True
parsed_origin = _origin_parts(origin)
parsed_host = _request_host_parts(request.host)
if parsed_origin is None or parsed_host is None:
return False
origin_host = parsed_origin.hostname.lower()
request_host = parsed_host.hostname.lower()
bind_host = os.environ.get("BRIDGE_HOST", "127.0.0.1").strip().lower()
if bind_host in _LOOPBACK_HOSTS:
return origin_host in _LOOPBACK_HOSTS and request_host in _LOOPBACK_HOSTS
return parsed_origin.netloc.lower() == request.host.lower()


def cors_headers(origin: str = ""):
if not origin:
return {}
return {
"Access-Control-Allow-Origin": "*",
"Access-Control-Allow-Methods": "GET,POST,PUT,DELETE,OPTIONS",
"Access-Control-Allow-Origin": origin,
"Access-Control-Allow-Methods": "GET,POST,PUT,PATCH,DELETE,OPTIONS",
"Access-Control-Allow-Headers": "Content-Type",
"Vary": "Origin",
}


@web.middleware
async def cors_middleware(request, handler):
origin = request.headers.get("Origin", "")
if origin and not _origin_matches_host(request, origin):
return web.json_response({"ok": False, "error": "cross-origin request rejected"}, status=403)
headers = cors_headers(origin)
if request.method == "OPTIONS":
return web.Response(status=204, headers=cors_headers())
return web.Response(status=204, headers=headers)
resp = await handler(request)
for k, v in cors_headers().items():
for k, v in headers.items():
resp.headers[k] = v
return resp


def json_ok(data: dict, status: int = 200):
return web.json_response(data, status=status, headers=cors_headers(), dumps=lambda x: json.dumps(x, ensure_ascii=False, default=str))
return web.json_response(data, status=status, dumps=lambda x: json.dumps(x, ensure_ascii=False, default=str))


async def read_json(request) -> dict:
Expand Down
116 changes: 116 additions & 0 deletions tests/test_desktop_bridge_origin.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,116 @@
import importlib.util
import os
import sys
import tempfile
import unittest
from pathlib import Path

from aiohttp import web
from aiohttp.test_utils import TestClient, TestServer


ROOT = Path(__file__).resolve().parents[1]
FRONTENDS = ROOT / "frontends"
if str(FRONTENDS) not in sys.path:
sys.path.insert(0, str(FRONTENDS))

# desktop_bridge has import-time persistence/upload setup. Point it at an isolated
# valid GA root so test discovery cannot touch a developer's configured GA_ROOT.
_TEST_GA_ROOT = tempfile.TemporaryDirectory()
_TEST_GA_PATH = Path(_TEST_GA_ROOT.name)
(_TEST_GA_PATH / "agentmain.py").touch()
_old_ga_root = os.environ.get("GA_ROOT")
_old_argv = sys.argv[:]
os.environ["GA_ROOT"] = str(_TEST_GA_PATH)
sys.argv = [sys.argv[0]]
try:
spec = importlib.util.spec_from_file_location("desktop_bridge", FRONTENDS / "desktop_bridge.py")
assert spec and spec.loader
bridge = importlib.util.module_from_spec(spec)
sys.modules[spec.name] = bridge
spec.loader.exec_module(bridge)
finally:
sys.argv = _old_argv
if _old_ga_root is None:
os.environ.pop("GA_ROOT", None)
else:
os.environ["GA_ROOT"] = _old_ga_root


class DesktopBridgeOriginTests(unittest.IsolatedAsyncioTestCase):
async def asyncSetUp(self):
app = web.Application(middlewares=[bridge.cors_middleware])

async def probe(request):
return bridge.json_ok({"ok": True})

app.router.add_post("/probe", probe)
self.client = TestClient(TestServer(app))
await self.client.start_server()

async def asyncTearDown(self):
await self.client.close()

async def test_import_uses_isolated_ga_root(self):
self.assertEqual(Path(bridge.DEFAULT_GA_ROOT), _TEST_GA_PATH)
self.assertEqual(bridge._WEB_UPLOAD_DIR, _TEST_GA_PATH / "temp" / "desktop_uploads")

async def test_cross_origin_preflight_is_rejected(self):
response = await self.client.options(
"/probe",
headers={
"Origin": "https://evil.example",
"Access-Control-Request-Method": "POST",
"Access-Control-Request-Headers": "content-type",
},
)
self.assertEqual(response.status, 403)

async def test_cross_origin_post_is_rejected(self):
response = await self.client.post(
"/probe", headers={"Origin": "https://evil.example"}, json={}
)
self.assertEqual(response.status, 403)

async def test_dns_rebinding_origin_is_rejected_on_loopback(self):
response = await self.client.post(
"/probe",
headers={"Host": "evil.example:14168", "Origin": "https://evil.example:14168"},
json={},
)
self.assertEqual(response.status, 403)

async def test_malformed_serialized_origins_are_rejected(self):
origin = str(self.client.make_url("/")).rstrip("/")
origins = (
f"{origin}/path",
f"{origin}?x=1",
origin.replace("http://", "http://user@", 1),
"http://127.0.0.1:notaport",
f"{origin}#fragment",
)
for candidate in origins:
with self.subTest(origin=candidate):
response = await self.client.post(
"/probe", headers={"Origin": candidate}, json={}
)
self.assertEqual(response.status, 403)

async def test_null_origin_is_rejected(self):
response = await self.client.post("/probe", headers={"Origin": "null"}, json={})
self.assertEqual(response.status, 403)

async def test_same_origin_request_is_allowed(self):
origin = str(self.client.make_url("/")).rstrip("/")
response = await self.client.post("/probe", headers={"Origin": origin}, json={})
self.assertEqual(response.status, 200)
self.assertEqual(response.headers.get("Access-Control-Allow-Origin"), origin)

async def test_non_browser_request_without_origin_is_allowed(self):
response = await self.client.post("/probe", json={})
self.assertEqual(response.status, 200)
self.assertNotIn("Access-Control-Allow-Origin", response.headers)


if __name__ == "__main__":
unittest.main()