Files
govoplan-calendar/tests/test_caldav_security.py

126 lines
4.8 KiB
Python

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()