# # This file is licensed under the Affero General Public License (AGPL) version 3. # # Copyright (C) 2026 Element Creations Ltd # # This program is free software: you can redistribute it and/or modify # it under the terms of the GNU Affero General Public License as # published by the Free Software Foundation, either version 3 of the # License, or (at your option) any later version. # # See the GNU Affero General Public License for more details: # . # import os import tempfile from unittest.mock import Mock import yaml from twisted.internet import defer from twisted.internet.testing import MemoryReactor from twisted.web.http_headers import Headers from synapse.rest import admin from synapse.rest.client import appservice_proxy, login from synapse.server import HomeServer from synapse.types import JsonDict from synapse.util.clock import Clock from synapse.util.json import json_encoder from tests import unittest from tests.test_utils import FakeResponse APPSERVICE_URL = "http://appservice.example.com" APPSERVICE_PREFIX = "rtc/livekit" VERSIONED_PREFIX = f"v1/{APPSERVICE_PREFIX}" class ApplicationServiceClientProxyTestCase(unittest.HomeserverTestCase): servlets = [ admin.register_servlets, login.register_servlets, appservice_proxy.register_servlets, ] def default_config(self) -> JsonDict: config = super().default_config() with tempfile.NamedTemporaryFile( mode="w", prefix="as_proxy_config", delete=False ) as f: self.addCleanup(os.remove, f.name) yaml.dump( { "id": "proxy_as", "url": None, "as_token": "as_token", "hs_token": "hs_token", "sender_localpart": "proxy_bot", "namespaces": {}, "io.element.msc4512.proxy_prefix": APPSERVICE_PREFIX, "io.element.msc4512.proxy_url": APPSERVICE_URL, }, f, ) config["app_service_config_files"] = [f.name] config.setdefault("experimental_features", {}).setdefault( "msc4512_enabled", True ) return config def prepare(self, _reactor: MemoryReactor, _clock: Clock, hs: HomeServer) -> None: self.agent = Mock() hs.get_proxied_http_client().agent = self.agent self.user_id = self.register_user("proxy_user", "password") self.access_token = self.login("proxy_user", "password") def test_get_is_proxied(self) -> None: self.agent.request = Mock( return_value=defer.succeed( FakeResponse.json(code=200, payload={"hello": "world"}) ) ) channel = self.make_request( "GET", f"/_matrix/client/{VERSIONED_PREFIX}/some/path?foo=bar", shorthand=False, access_token=self.access_token, ) self.assertEqual(channel.code, 200) self.assertEqual(channel.json_body, {"hello": "world"}) ((method, uri), kwargs) = self.agent.request.call_args self.assertEqual(method, b"GET") self.assertEqual( uri, f"{APPSERVICE_URL}/_matrix/client/{VERSIONED_PREFIX}/some/path?foo=bar".encode(), ) headers: Headers = kwargs["headers"] self.assertEqual(headers.getRawHeaders(b"Authorization"), [b"Bearer hs_token"]) self.assertEqual( headers.getRawHeaders(b"X-Matrix-User-Identifier"), [self.user_id.encode("ascii")], ) def test_access_token_query_param_is_stripped(self) -> None: self.agent.request = Mock( return_value=defer.succeed( FakeResponse.json(code=200, payload={"hello": "world"}) ) ) channel = self.make_request( "GET", f"/_matrix/client/{VERSIONED_PREFIX}/some/path?access_token={self.access_token}&foo=bar", shorthand=False, ) self.assertEqual(channel.code, 200) self.assertEqual(channel.json_body, {"hello": "world"}) ((method, uri), _kwargs) = self.agent.request.call_args self.assertEqual(method, b"GET") self.assertEqual( uri, f"{APPSERVICE_URL}/_matrix/client/{VERSIONED_PREFIX}/some/path?foo=bar".encode(), ) def test_get_is_proxied_at_root_path(self) -> None: self.agent.request = Mock( return_value=defer.succeed( FakeResponse.json(code=200, payload={"hello": "world"}) ) ) channel = self.make_request( "GET", f"/_matrix/client/{VERSIONED_PREFIX}", shorthand=False, access_token=self.access_token, ) self.assertEqual(channel.code, 200) self.assertEqual(channel.json_body, {"hello": "world"}) ((method, uri), kwargs) = self.agent.request.call_args self.assertEqual(method, b"GET") self.assertEqual( uri, f"{APPSERVICE_URL}/_matrix/client/{VERSIONED_PREFIX}".encode(), ) headers: Headers = kwargs["headers"] self.assertEqual(headers.getRawHeaders(b"Authorization"), [b"Bearer hs_token"]) self.assertEqual( headers.getRawHeaders(b"X-Matrix-User-Identifier"), [self.user_id.encode("ascii")], ) def test_post_is_proxied(self) -> None: self.agent.request = Mock( return_value=defer.succeed( FakeResponse.json(code=200, payload={"hello": "world"}) ) ) channel = self.make_request( "POST", f"/_matrix/client/{VERSIONED_PREFIX}/some/path", content={"key": "value"}, shorthand=False, access_token=self.access_token, ) self.assertEqual(channel.code, 200) self.assertEqual(channel.json_body, {"hello": "world"}) ((method, uri), kwargs) = self.agent.request.call_args self.assertEqual(method, b"POST") self.assertEqual( uri, f"{APPSERVICE_URL}/_matrix/client/{VERSIONED_PREFIX}/some/path".encode(), ) headers: Headers = kwargs["headers"] self.assertEqual(headers.getRawHeaders(b"Authorization"), [b"Bearer hs_token"]) self.assertEqual( headers.getRawHeaders(b"X-Matrix-User-Identifier"), [self.user_id.encode("ascii")], ) self.assertEqual(headers.getRawHeaders(b"Content-Type"), [b"application/json"]) body_producer = kwargs["bodyProducer"] expected_body = json_encoder.encode({"key": "value"}).encode("utf8") self.assertEqual(body_producer.length, len(expected_body)) def test_headers_outside_the_allowlist_not_forwarded(self) -> None: self.agent.request = Mock( return_value=defer.succeed(FakeResponse.json(code=200, payload={})) ) self.make_request( "GET", f"/_matrix/client/{VERSIONED_PREFIX}/some/path", shorthand=False, access_token=self.access_token, custom_headers=[("Connection", "close"), ("X-Forward", "forward")], ) ((_method, _uri), kwargs) = self.agent.request.call_args headers: Headers = kwargs["headers"] self.assertIsNone(headers.getRawHeaders(b"Connection")) self.assertIsNone(headers.getRawHeaders(b"X-Forward")) def test_allowlisted_headers_forwarded(self) -> None: self.agent.request = Mock( return_value=defer.succeed(FakeResponse.json(code=200, payload={})) ) self.make_request( "GET", f"/_matrix/client/{VERSIONED_PREFIX}/some/path", shorthand=False, access_token=self.access_token, custom_headers=[ ("Accept", "application/json"), ("Accept-Language", "en-US"), ], ) ((_method, _uri), kwargs) = self.agent.request.call_args headers: Headers = kwargs["headers"] self.assertEqual(headers.getRawHeaders(b"Accept"), [b"application/json"]) self.assertEqual(headers.getRawHeaders(b"Accept-Language"), [b"en-US"]) def test_host_and_content_length_headers_not_forwarded(self) -> None: self.agent.request = Mock( return_value=defer.succeed(FakeResponse.json(code=200, payload={})) ) self.make_request( "POST", f"/_matrix/client/{VERSIONED_PREFIX}/some/path", content={"key": "value"}, shorthand=False, access_token=self.access_token, custom_headers=[("Host", "original-client-facing-host.example")], ) ((_method, _uri), kwargs) = self.agent.request.call_args headers: Headers = kwargs["headers"] self.assertIsNone(headers.getRawHeaders(b"Host")) self.assertIsNone(headers.getRawHeaders(b"Content-Length")) def test_response_headers_forwarded(self) -> None: self.agent.request = Mock( return_value=defer.succeed( FakeResponse( code=200, body=b"hello", headers=Headers({"X-Forward": ["forward"]}), ) ) ) channel = self.make_request( "GET", f"/_matrix/client/{VERSIONED_PREFIX}/some/path", shorthand=False, access_token=self.access_token, ) self.assertEqual(channel.code, 200) self.assertEqual(channel.result["body"], b"hello") self.assertEqual(channel.headers.getRawHeaders(b"X-Forward"), [b"forward"]) def test_response_cors_headers_set(self) -> None: self.agent.request = Mock( return_value=defer.succeed(FakeResponse.json(code=200, payload={})) ) channel = self.make_request( "GET", f"/_matrix/client/{VERSIONED_PREFIX}/some/path", shorthand=False, access_token=self.access_token, ) self.assertEqual(channel.code, 200) self.assertEqual( channel.headers.getRawHeaders(b"Access-Control-Allow-Origin"), [b"*"] ) def test_non_existing_path_under_proxy_prefix_is_rejected(self) -> None: self.agent.request = Mock(return_value=defer.fail(Exception("boom"))) channel = self.make_request( "GET", f"/_matrix/client/{VERSIONED_PREFIX}/does/not/exist", shorthand=False, access_token=self.access_token, ) self.assertEqual(channel.code, 500) self.assertEqual(channel.json_body["errcode"], "M_UNKNOWN") self.agent.request.assert_called() def test_unauthenticated_get_is_rejected(self) -> None: self.agent.request = Mock( return_value=defer.succeed(FakeResponse.json(code=200, payload={})) ) channel = self.make_request( "GET", f"/_matrix/client/{VERSIONED_PREFIX}/some/path", shorthand=False, ) self.assertEqual(channel.code, 401) self.agent.request.assert_not_called() @unittest.override_config({"rc_message": {"burst_count": 0}}) def test_rate_limited_request_is_rejected(self) -> None: self.agent.request = Mock( return_value=defer.succeed(FakeResponse.json(code=200, payload={})) ) channel = self.make_request( "GET", f"/_matrix/client/{VERSIONED_PREFIX}/some/path", shorthand=False, access_token=self.access_token, ) self.assertEqual(channel.code, 429) self.agent.request.assert_not_called() def test_path_with_dot_segment_is_rejected(self) -> None: self.agent.request = Mock( return_value=defer.succeed(FakeResponse.json(code=200, payload={})) ) channel = self.make_request( "GET", f"/_matrix/client/{VERSIONED_PREFIX}/some/../path", shorthand=False, access_token=self.access_token, ) self.assertEqual(channel.code, 400) self.assertEqual(channel.json_body["errcode"], "M_INVALID_PARAM") self.agent.request.assert_not_called() def test_path_with_encoded_dot_segment_is_rejected(self) -> None: self.agent.request = Mock( return_value=defer.succeed(FakeResponse.json(code=200, payload={})) ) channel = self.make_request( "GET", f"/_matrix/client/{VERSIONED_PREFIX}/some/%2e%2e/path", shorthand=False, access_token=self.access_token, ) self.assertEqual(channel.code, 400) self.assertEqual(channel.json_body["errcode"], "M_INVALID_PARAM") self.agent.request.assert_not_called() def test_unregistered_prefix_is_rejected(self) -> None: channel = self.make_request( "GET", "/_matrix/client/not-a-prefix", shorthand=False, ) self.assertEqual(channel.code, 404) def test_unregistered_prefix_with_suffix_is_rejected(self) -> None: channel = self.make_request( "GET", f"/_matrix/client/{VERSIONED_PREFIX}-2", shorthand=False, ) self.assertEqual(channel.code, 404) def test_missing_version_segment_is_rejected(self) -> None: channel = self.make_request( "GET", f"/_matrix/client/{APPSERVICE_PREFIX}/some/path", shorthand=False, access_token=self.access_token, ) self.assertEqual(channel.code, 404) self.agent.request.assert_not_called() def test_unstable_msc_version_segment_is_proxied(self) -> None: self.agent.request = Mock( return_value=defer.succeed( FakeResponse.json(code=200, payload={"hello": "world"}) ) ) path = f"/_matrix/client/unstable/org.example.msc9999/{APPSERVICE_PREFIX}/some/path" channel = self.make_request( "GET", path, shorthand=False, access_token=self.access_token, ) self.assertEqual(channel.code, 200) self.assertEqual(channel.json_body, {"hello": "world"}) ((method, uri), _kwargs) = self.agent.request.call_args self.assertEqual(method, b"GET") self.assertEqual(uri, f"{APPSERVICE_URL}{path}".encode()) @unittest.override_config({"experimental_features": {"msc4512_enabled": False}}) def test_proxy_route_not_registered_when_msc4512_disabled(self) -> None: channel = self.make_request( "GET", f"/_matrix/client/{VERSIONED_PREFIX}/some/path", shorthand=False, access_token=self.access_token, ) self.assertEqual(channel.code, 404) self.agent.request.assert_not_called()