from __future__ import annotations import contextlib import threading import unittest from collections.abc import Iterator from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from govoplan_calendar.backend.caldav import ( CalDAVClient, CalDAVError, absolute_dav_url, urllib_transport, ) @contextlib.contextmanager def running_http_server(handler: type[BaseHTTPRequestHandler]) -> Iterator[str]: server = ThreadingHTTPServer(("127.0.0.1", 0), handler) thread = threading.Thread(target=server.serve_forever, daemon=True) thread.start() try: host, port = server.server_address yield f"http://{host}:{port}" finally: server.shutdown() server.server_close() thread.join(timeout=2) class CalDAVUrlSecurityTests(unittest.TestCase): def test_discovery_href_must_remain_on_configured_origin(self) -> None: base_url = "https://dav.example.test/calendars/ada/" self.assertEqual( absolute_dav_url(base_url, "/principals/users/ada/"), "https://dav.example.test/principals/users/ada/", ) with self.assertRaisesRegex(CalDAVError, "configured collection origin"): absolute_dav_url(base_url, "https://evil.example.test/steal/") with self.assertRaisesRegex(CalDAVError, "query or fragment"): absolute_dav_url(base_url, "/principals/users/ada/?token=secret") def test_object_href_rejects_userinfo_query_and_fragment(self) -> None: client = CalDAVClient(collection_url="https://dav.example.test/calendars/ada") with self.assertRaisesRegex(CalDAVError, "embedded credentials"): client.object_url("https://user:secret@dav.example.test/calendars/ada/event.ics") with self.assertRaisesRegex(CalDAVError, "query or fragment"): client.object_url("/calendars/ada/event.ics?download=1") with self.assertRaisesRegex(CalDAVError, "query or fragment"): client.object_url("/calendars/ada/event.ics#fragment") self.assertEqual( client.object_url("/calendars/ada/event.ics"), "https://dav.example.test/calendars/ada/event.ics", ) def test_transport_refuses_redirect_before_forwarding_authorization(self) -> None: forwarded_authorization: list[str | None] = [] class TargetHandler(BaseHTTPRequestHandler): def do_GET(self) -> None: # noqa: N802 - BaseHTTPRequestHandler API forwarded_authorization.append(self.headers.get("Authorization")) self.send_response(200) self.end_headers() def log_message(self, _format: str, *_args: object) -> None: return with running_http_server(TargetHandler) as target_url: class RedirectHandler(BaseHTTPRequestHandler): def do_GET(self) -> None: # noqa: N802 - BaseHTTPRequestHandler API self.send_response(302) self.send_header("Location", f"{target_url}/stolen.ics") self.end_headers() def log_message(self, _format: str, *_args: object) -> None: return with running_http_server(RedirectHandler) as redirect_url: status, _headers, _body = urllib_transport( "GET", f"{redirect_url}/event.ics", {"Authorization": "Bearer top-secret"}, None, 2, ) self.assertEqual(status, 302) self.assertEqual(forwarded_authorization, []) def test_transport_preserves_same_origin_redirects(self) -> None: forwarded_authorization: list[str | None] = [] class RedirectHandler(BaseHTTPRequestHandler): def do_GET(self) -> None: # noqa: N802 - BaseHTTPRequestHandler API if self.path == "/event.ics": self.send_response(302) self.send_header("Location", "/redirected.ics") self.end_headers() return forwarded_authorization.append(self.headers.get("Authorization")) self.send_response(200) self.end_headers() self.wfile.write(b"calendar") def log_message(self, _format: str, *_args: object) -> None: return with running_http_server(RedirectHandler) as source_url: status, _headers, body = urllib_transport( "GET", f"{source_url}/event.ics", {"Authorization": "Bearer expected"}, None, 2, ) self.assertEqual(status, 200) self.assertEqual(body, b"calendar") self.assertEqual(forwarded_authorization, ["Bearer expected"]) if __name__ == "__main__": unittest.main()