188 lines
10 KiB
Python
188 lines
10 KiB
Python
"""Origin checks use operator configuration, never caller-controlled routing headers."""
|
|
|
|
from pathlib import Path
|
|
import tempfile
|
|
from types import SimpleNamespace
|
|
import unittest
|
|
from unittest.mock import AsyncMock, patch
|
|
|
|
from fastapi.testclient import TestClient
|
|
|
|
from backend.app import db, main
|
|
from backend.app.config import settings
|
|
from backend.app.routers import auth as auth_router
|
|
from backend.app.services import public_urls
|
|
from backend.app.services.request_origins import is_allowed_request_origin
|
|
|
|
|
|
PUBLIC_ORIGIN = "https://watch.example.test"
|
|
LOCAL_ORIGIN = "http://localhost:3000"
|
|
|
|
|
|
class RequestOriginTests(unittest.TestCase):
|
|
def setUp(self):
|
|
self.runtime = SimpleNamespace(
|
|
magent_proxy_enabled=False,
|
|
magent_proxy_base_url=None,
|
|
magent_application_url=PUBLIC_ORIGIN,
|
|
)
|
|
self.enterContext(patch.object(settings, "cors_allow_origin", LOCAL_ORIGIN))
|
|
self.enterContext(patch.object(public_urls, "get_runtime_settings", return_value=self.runtime))
|
|
|
|
def test_explicit_cors_and_configured_public_url_are_both_allowed(self):
|
|
self.assertTrue(is_allowed_request_origin(LOCAL_ORIGIN))
|
|
self.assertTrue(is_allowed_request_origin(PUBLIC_ORIGIN))
|
|
self.assertFalse(is_allowed_request_origin("https://unrelated.example.test"))
|
|
|
|
def test_scheme_hostname_case_and_default_ports_are_canonicalized(self):
|
|
for origin in (PUBLIC_ORIGIN, "HTTPS://WATCH.EXAMPLE.TEST", "https://watch.example.test:443"):
|
|
with self.subTest(origin=origin):
|
|
self.assertTrue(is_allowed_request_origin(origin))
|
|
self.runtime.magent_application_url = "http://watch.example.test:80"
|
|
self.assertTrue(is_allowed_request_origin("http://WATCH.example.test"))
|
|
self.assertFalse(is_allowed_request_origin("https://watch.example.test"))
|
|
self.assertFalse(is_allowed_request_origin("http://watch.example.test:8080"))
|
|
|
|
def test_nondefault_ports_must_match(self):
|
|
self.runtime.magent_application_url = "https://watch.example.test:8443/magent"
|
|
self.assertTrue(is_allowed_request_origin("https://watch.example.test:8443"))
|
|
self.assertFalse(is_allowed_request_origin("https://watch.example.test"))
|
|
self.assertFalse(is_allowed_request_origin("https://watch.example.test:443"))
|
|
|
|
def test_configured_subpath_does_not_become_part_of_origin(self):
|
|
self.runtime.magent_application_url = PUBLIC_ORIGIN + "/magent/"
|
|
self.assertTrue(is_allowed_request_origin(PUBLIC_ORIGIN))
|
|
self.assertFalse(is_allowed_request_origin(PUBLIC_ORIGIN + "/magent"))
|
|
|
|
def test_enabled_proxy_uses_configured_proxy_public_url(self):
|
|
self.runtime.magent_proxy_enabled = True
|
|
self.runtime.magent_proxy_base_url = "https://proxy.example.test/magent"
|
|
self.assertTrue(is_allowed_request_origin("https://proxy.example.test"))
|
|
self.assertTrue(is_allowed_request_origin(LOCAL_ORIGIN))
|
|
self.assertFalse(is_allowed_request_origin(PUBLIC_ORIGIN))
|
|
|
|
def test_unconfigured_public_url_only_allows_explicit_cors(self):
|
|
self.runtime.magent_application_url = None
|
|
self.assertTrue(is_allowed_request_origin(LOCAL_ORIGIN))
|
|
self.assertFalse(is_allowed_request_origin(PUBLIC_ORIGIN))
|
|
|
|
def test_invalid_or_non_origin_inputs_are_rejected(self):
|
|
for origin in (
|
|
"", "null", "*", "watch.example.test", "//watch.example.test",
|
|
"ftp://watch.example.test", "javascript:alert(1)",
|
|
PUBLIC_ORIGIN + "/", PUBLIC_ORIGIN + "/path",
|
|
PUBLIC_ORIGIN + "?query=true", PUBLIC_ORIGIN + "#fragment",
|
|
PUBLIC_ORIGIN + "?", PUBLIC_ORIGIN + "#",
|
|
"https://user@watch.example.test", "https://user:password@watch.example.test",
|
|
"https://watch.example.test@evil.example.test", "https://watch.example.test.evil.example.test",
|
|
"https://watch.example.test:0", "https://watch.example.test:65536",
|
|
"https://watch.example.test:invalid", "https://[invalid",
|
|
PUBLIC_ORIGIN + " https://evil.example.test", PUBLIC_ORIGIN + ",https://evil.example.test",
|
|
"https://watch.example.test\\@evil.example.test", PUBLIC_ORIGIN + "\n",
|
|
):
|
|
with self.subTest(origin=repr(origin)):
|
|
self.assertFalse(is_allowed_request_origin(origin))
|
|
|
|
def test_invalid_configured_public_url_does_not_authorize_an_origin(self):
|
|
for configured in (
|
|
"https://user:password@watch.example.test", PUBLIC_ORIGIN + "?token=private",
|
|
PUBLIC_ORIGIN + "#fragment", "javascript:alert(1)", "https://watch.example.test:65536",
|
|
):
|
|
with self.subTest(configured=configured):
|
|
self.runtime.magent_application_url = configured
|
|
self.assertFalse(is_allowed_request_origin(PUBLIC_ORIGIN))
|
|
self.assertTrue(is_allowed_request_origin(LOCAL_ORIGIN))
|
|
|
|
|
|
class RequestOriginHttpTests(unittest.TestCase):
|
|
def setUp(self):
|
|
temporary = tempfile.TemporaryDirectory(ignore_cleanup_errors=True)
|
|
self.addCleanup(temporary.cleanup)
|
|
self.runtime = SimpleNamespace(
|
|
magent_proxy_enabled=False,
|
|
magent_proxy_base_url=None,
|
|
magent_application_url=PUBLIC_ORIGIN,
|
|
)
|
|
for name, value in {
|
|
"sqlite_path": str(Path(temporary.name) / "origin-tests.db"),
|
|
"sqlite_journal_mode": "DELETE",
|
|
"jwt_secret": "request-origin-tests-jwt-secret-at-least-32-characters",
|
|
"settings_encryption_key": None,
|
|
"admin_username": "unused-environment-admin",
|
|
"admin_password": "",
|
|
"cors_allow_origin": LOCAL_ORIGIN,
|
|
"auth_cookie_domain": None,
|
|
"auth_cookie_secure": True,
|
|
}.items():
|
|
self.enterContext(patch.object(settings, name, value))
|
|
self.enterContext(patch.object(public_urls, "get_runtime_settings", return_value=self.runtime))
|
|
# Constructing without a context deliberately skips production startup:
|
|
# no migrations/workers/listeners/log files outside this temporary DB.
|
|
db.init_db()
|
|
self.client = TestClient(main.app, base_url=PUBLIC_ORIGIN)
|
|
self.addCleanup(self.client.close)
|
|
|
|
def test_public_origin_reaches_both_auth_handlers_with_localhost_cors_default(self):
|
|
for path in ("/auth/login", "/auth/jellyfin/login"):
|
|
with self.subTest(path=path):
|
|
response = self.client.post(path, data={}, headers={"Origin": PUBLIC_ORIGIN})
|
|
self.assertEqual(response.status_code, 422, response.text)
|
|
self.assertNotEqual(response.json().get("detail"), "Cross-origin state change rejected")
|
|
|
|
def test_explicit_cors_origin_remains_allowed(self):
|
|
response = self.client.post("/auth/login", data={}, headers={"Origin": LOCAL_ORIGIN})
|
|
self.assertEqual(response.status_code, 422, response.text)
|
|
|
|
def test_no_origin_keeps_existing_nonbrowser_behavior(self):
|
|
response = self.client.post("/auth/login", data={})
|
|
self.assertEqual(response.status_code, 422, response.text)
|
|
|
|
def test_caller_controlled_host_forwarding_and_fetch_headers_cannot_authorize_evil_origin(self):
|
|
for path in ("/auth/login", "/auth/jellyfin/login"):
|
|
for routing_headers in (
|
|
{},
|
|
{"Host": "evil.example.test"},
|
|
{"X-Forwarded-Host": "evil.example.test", "X-Forwarded-Proto": "https"},
|
|
{"Host": "evil.example.test", "X-Forwarded-Host": "evil.example.test", "Sec-Fetch-Site": "same-origin"},
|
|
{"Host": "watch.example.test", "X-Forwarded-Host": "watch.example.test", "Sec-Fetch-Site": "same-origin"},
|
|
):
|
|
with self.subTest(path=path, routing_headers=routing_headers):
|
|
response = self.client.post(path, data={}, headers={"Origin": "https://evil.example.test", **routing_headers})
|
|
self.assertEqual(response.status_code, 403, response.text)
|
|
self.assertEqual(response.json()["detail"], "Cross-origin state change rejected")
|
|
|
|
def test_null_path_query_and_userinfo_origins_are_rejected_before_login(self):
|
|
for origin in ("null", PUBLIC_ORIGIN + "/", PUBLIC_ORIGIN + "/path", PUBLIC_ORIGIN + "?query=1", "https://user@watch.example.test"):
|
|
with self.subTest(origin=origin):
|
|
response = self.client.post("/auth/login", data={}, headers={"Origin": origin})
|
|
self.assertEqual(response.status_code, 403, response.text)
|
|
|
|
def test_valid_local_login_works_from_configured_public_origin(self):
|
|
password = "origin-tests-valid-local-password"
|
|
db.create_user("origin-owner", password, role="admin")
|
|
response = self.client.post("/auth/login", data={"username": "origin-owner", "password": password}, headers={"Origin": PUBLIC_ORIGIN})
|
|
self.assertEqual(response.status_code, 200, response.text)
|
|
self.assertIn(settings.auth_cookie_name, self.client.cookies)
|
|
profile = self.client.get("/auth/profile")
|
|
self.assertEqual(profile.status_code, 200, profile.text)
|
|
self.assertEqual(profile.json()["user"]["username"], "origin-owner")
|
|
|
|
def test_valid_mocked_jellyfin_login_works_from_configured_public_origin(self):
|
|
jellyfin_runtime = SimpleNamespace(jellyfin_base_url="http://jellyfin.test:8096", jellyfin_api_key="test-api-key")
|
|
upstream = SimpleNamespace(
|
|
configured=lambda: True,
|
|
authenticate_by_name=AsyncMock(return_value={"User": {"Id": "test-jellyfin-id", "Name": "origin-viewer"}}),
|
|
get_users=AsyncMock(return_value=[]),
|
|
_extract_user_id=lambda _response: "test-jellyfin-id",
|
|
)
|
|
with patch.object(auth_router, "get_runtime_settings", return_value=jellyfin_runtime), patch.object(auth_router, "JellyfinClient", return_value=upstream), patch.object(auth_router, "get_cached_jellyseerr_users", return_value=[]):
|
|
response = self.client.post("/auth/jellyfin/login", data={"username": "origin-viewer", "password": "origin-tests-jellyfin-password"}, headers={"Origin": PUBLIC_ORIGIN})
|
|
self.assertEqual(response.status_code, 200, response.text)
|
|
upstream.authenticate_by_name.assert_awaited_once_with("origin-viewer", "origin-tests-jellyfin-password")
|
|
self.assertIn(settings.auth_cookie_name, self.client.cookies)
|
|
self.assertEqual(db.get_user_by_username("origin-viewer")["auth_provider"], "jellyfin")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|