diff --git a/changelog.d/19972.feature b/changelog.d/19972.feature new file mode 100644 index 0000000000..56e3b27136 --- /dev/null +++ b/changelog.d/19972.feature @@ -0,0 +1 @@ +Add experimental support for letting application services proxy namespaces in the C-S and S-S API as per MSC4512. diff --git a/synapse/appservice/__init__.py b/synapse/appservice/__init__.py index 19fff7f009..20f7a410b6 100644 --- a/synapse/appservice/__init__.py +++ b/synapse/appservice/__init__.py @@ -97,6 +97,11 @@ class ApplicationService: # values. NS_LIST = [NS_USERS, NS_ALIASES, NS_ROOMS] + # Prefixes are applied after the version segment(s) (either /vX/ or /unstable/foo/): + # - /_matrix/client/(unstable/[^/]+|v[^/]+)/{prefix}/.* + # - /_matrix/federation/(unstable/[^/]+|v[^/]+)/{prefix}/.* + ALLOWED_PROXY_PREFIXES = {"rtc/livekit"} + def __init__( self, token: str, @@ -113,11 +118,19 @@ class ApplicationService: msc3202_transaction_extensions: bool = False, msc4190_device_management: bool = False, scopes: Iterable[str] = frozenset(), + proxy_prefix: str | None = None, + proxy_url: str | None = None, ): self.token = token self.url = ( url.rstrip("/") if isinstance(url, str) else None ) # url must not end with a slash + self.proxy_url = ( + proxy_url.rstrip("/") if isinstance(proxy_url, str) else None + ) # proxy_url must not end with a slash + self.proxy_prefix = ( + proxy_prefix.rstrip("/") if isinstance(proxy_prefix, str) else None + ) # proxy_prefix must not end with a slash self.hs_token = hs_token # The full Matrix ID for this application service's sender. self.sender = sender @@ -143,6 +156,14 @@ class ApplicationService: if "|" in self.id: raise Exception("application service ID cannot contain '|' character") + if (self.proxy_prefix is None) != (self.proxy_url is None): + raise KeyError("proxy_url and proxy_prefix must always be set together") + if proxy_prefix is not None: + if not proxy_prefix or not self.proxy_url: + raise ValueError("proxy_prefix and proxy_url must be non-empty strings") + if not self._is_proxy_prefix_allowed(proxy_prefix): + raise ValueError(f"cannot claim reserved proxy prefix {proxy_prefix!r}") + # .protocols is a publicly visible field if protocols: self.protocols = set(protocols) @@ -206,6 +227,12 @@ class ApplicationService: return namespace.exclusive return False + def _is_proxy_prefix_allowed(self, prefix: str) -> bool: + return any( + prefix == allowed or prefix.startswith(allowed + "/") + for allowed in ApplicationService.ALLOWED_PROXY_PREFIXES + ) + @cached(num_args=1, cache_context=True) async def _matches_user_in_member_list( self, diff --git a/synapse/config/appservice.py b/synapse/config/appservice.py index 4e61ef694f..1868985f80 100644 --- a/synapse/config/appservice.py +++ b/synapse/config/appservice.py @@ -68,6 +68,7 @@ def load_appservices( # Dicts of value -> filename seen_as_tokens: dict[str, str] = {} seen_ids: dict[str, str] = {} + seen_proxy_prefixes: dict[str, str] = {} appservices = [] @@ -93,6 +94,17 @@ def load_appservices( ) ) seen_as_tokens[appservice.token] = config_file + if appservice.proxy_prefix is not None: + for seen_prefix, seen_file in seen_proxy_prefixes.items(): + if _proxy_prefixes_overlap( + appservice.proxy_prefix, seen_prefix + ): + raise ConfigError( + "io.element.msc4512.proxy_prefix values must not overlap across " + "application services: " + f"{appservice.proxy_prefix} (files: {config_file}, {seen_file})" + ) + seen_proxy_prefixes[appservice.proxy_prefix] = config_file logger.info("Loaded application service: %s", appservice) appservices.append(appservice) except Exception as e: @@ -102,6 +114,15 @@ def load_appservices( return appservices +def _proxy_prefixes_overlap(prefix_a: str, prefix_b: str) -> bool: + """Returns whether two proxy prefixes overlap by sharing a common path prefix.""" + return ( + prefix_a == prefix_b + or prefix_a.startswith(prefix_b + "/") + or prefix_b.startswith(prefix_a + "/") + ) + + def _load_appservice( hostname: str, as_info: JsonDict, config_filename: str ) -> ApplicationService: @@ -207,6 +228,22 @@ def _load_appservice( "The `io.element.msc4502.scopes` option should be a list of strings if specified." ) + # Opt-in setting to enable proxying C-S and S-S API endpoints. + # When set, Synapse will reverse-proxy requests under the prefix to the appservice: + proxy_prefix = as_info.get("io.element.msc4512.proxy_prefix") + if proxy_prefix is not None: + if not isinstance(proxy_prefix, str) or not proxy_prefix: + raise ValueError( + "The `io.element.msc4512.proxy_prefix` option should be a non-empty string." + ) + + proxy_url = as_info.get("io.element.msc4512.proxy_url") + if proxy_url is not None: + if not isinstance(proxy_url, str) or not proxy_url: + raise ValueError( + "The `io.element.msc4512.proxy_url` option should be a non-empty string." + ) + return ApplicationService( token=as_info["as_token"], url=as_info["url"], @@ -222,4 +259,6 @@ def _load_appservice( msc3202_transaction_extensions=msc3202_transaction_extensions, msc4190_device_management=msc4190_enabled, scopes=scopes, + proxy_prefix=proxy_prefix, + proxy_url=proxy_url, ) diff --git a/synapse/config/experimental.py b/synapse/config/experimental.py index 97dc803ab7..1c2f021322 100644 --- a/synapse/config/experimental.py +++ b/synapse/config/experimental.py @@ -312,3 +312,6 @@ class ExperimentalConfig(Config): # MSC4491: Invite reasons in room creation self.msc4491_enabled: bool = experimental.get("msc4491_enabled", False) + + # MSC4512: Delegating parts of the C-S and S-S API to application services + self.msc4512_enabled: bool = experimental.get("msc4512_enabled", False) diff --git a/synapse/federation/transport/server/__init__.py b/synapse/federation/transport/server/__init__.py index 0eff49cf73..70fc55123b 100644 --- a/synapse/federation/transport/server/__init__.py +++ b/synapse/federation/transport/server/__init__.py @@ -24,6 +24,7 @@ import logging from typing import TYPE_CHECKING, Iterable, Literal from synapse.api.errors import FederationDeniedError, SynapseError +from synapse.federation.transport.server import appservice_proxy from synapse.federation.transport.server._base import ( Authenticator, BaseFederationServlet, @@ -340,3 +341,6 @@ def register_servlets( ratelimiter=ratelimiter, server_name=hs.hostname, ).register(resource) + + if "federation" in servlet_groups: + appservice_proxy.register_servlets(hs, resource, authenticator, ratelimiter) diff --git a/synapse/federation/transport/server/appservice_proxy.py b/synapse/federation/transport/server/appservice_proxy.py new file mode 100644 index 0000000000..7ececd8fff --- /dev/null +++ b/synapse/federation/transport/server/appservice_proxy.py @@ -0,0 +1,115 @@ +# +# 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 logging +import re +from http import HTTPStatus +from io import BytesIO +from typing import TYPE_CHECKING + +from synapse.api.errors import Codes, SynapseError +from synapse.appservice import ApplicationService +from synapse.federation.transport.server._base import Authenticator +from synapse.http import QuieterFileBodyProducer +from synapse.http.appservice_proxy import proxy_request_to_appservice +from synapse.http.server import HttpServer, ServletCallback +from synapse.http.site import SynapseRequest +from synapse.util.json import json_decoder +from synapse.util.ratelimitutils import FederationRateLimiter + +if TYPE_CHECKING: + from synapse.server import HomeServer + +logger = logging.getLogger(__name__) + + +def _make_proxy_callback( + hs: "HomeServer", + authenticator: Authenticator, + ratelimiter: FederationRateLimiter, + appservice: ApplicationService, +) -> ServletCallback: + """Builds a servlet callback that authenticates an incoming federation request, + rate-limits it by origin, and forwards it to the given application service's + proxy URL. + """ + + async def _proxy(request: SynapseRequest, **kwargs: str) -> None: + body_producer = None + content = None + + if request.method in (b"PUT", b"POST"): + raw_body = request.content.read() # type: ignore[union-attr] + try: + content = json_decoder.decode(raw_body.decode("utf-8")) + except Exception: + raise SynapseError( + HTTPStatus.BAD_REQUEST, "Content not JSON.", Codes.NOT_JSON + ) + body_producer = QuieterFileBodyProducer(BytesIO(raw_body)) + + origin = await authenticator.authenticate_request(request, content) + + # Apply the same per-origin rate limiting that every other federation endpoint gets. + with ratelimiter.ratelimit(origin) as d: + await d + if request._disconnected: + logger.warning( + "client disconnected before we started processing request" + ) + return + + await proxy_request_to_appservice( + request, + hs, + appservice, + body_producer, + extra_request_headers={b"X-Matrix-Origin": origin.encode("ascii")}, + ) + + return _proxy + + +def register_servlets( + hs: "HomeServer", + resource: HttpServer, + authenticator: Authenticator, + ratelimiter: FederationRateLimiter, +) -> None: + """Registers blanket reverse-proxy routes for each application service that has + configured a proxy prefix. This forwards requests under + /_matrix/federation///* (where is either "vN" or "unstable") + to the same path under the application service's proxy URL after verifying request + authentication. + """ + if not hs.config.experimental.msc4512_enabled: + return + + for appservice in hs.get_datastores().main.get_app_services(): + if appservice.proxy_prefix is None or appservice.proxy_url is None: + continue + + pattern = re.compile( + r"^/_matrix/federation/(?:unstable/[^/]+|v[^/]+)/%s(/.*)?$" + % (re.escape(appservice.proxy_prefix),) + ) + callback = _make_proxy_callback(hs, authenticator, ratelimiter, appservice) + + for method in ("GET", "POST", "PUT", "DELETE"): + resource.register_paths( + method, + (pattern,), + callback, + "ApplicationServiceFederationProxy", + ) diff --git a/synapse/http/appservice_proxy.py b/synapse/http/appservice_proxy.py new file mode 100644 index 0000000000..714dfcc0c1 --- /dev/null +++ b/synapse/http/appservice_proxy.py @@ -0,0 +1,194 @@ +# +# 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 logging +from http import HTTPStatus +from typing import TYPE_CHECKING, Optional, cast +from urllib.parse import parse_qs, unquote_to_bytes, urlencode, urlsplit + +from twisted.python import failure +from twisted.web.http_headers import Headers +from twisted.web.iweb import IBodyProducer, IResponse + +from synapse.api.errors import Codes, SynapseError +from synapse.appservice import ApplicationService +from synapse.http.proxy import ( + HOP_BY_HOP_HEADERS_LOWERCASE, + _ProxyResponseBody, + parse_connection_header_value, +) +from synapse.http.server import return_json_error, set_cors_headers +from synapse.http.site import SynapseRequest +from synapse.logging.context import make_deferred_yieldable, run_in_background +from synapse.util.async_helpers import timeout_deferred + +if TYPE_CHECKING: + from synapse.server import HomeServer + +logger = logging.getLogger(__name__) + + +def has_dot_segments(path: bytes) -> bool: + """Whether the given request path contains any "." or ".." segments. + + The path is percent-decoded before it is split, since `%2e%2e` and `..` are + equivalent to anything that resolves the path, and `%2f` hides a separator that + would otherwise not be seen. Note that a single decode is deliberate: it matches + the single decode that route arguments get in `JsonResource._async_render`, so + `%252e%252e` is left alone rather than being treated as a dot segment. + + The caller is expected to reject such a path rather than rewrite it. Synapse + routes on the raw path, and federation request signatures cover the raw URI, so + normalising a path in place would break both. + """ + return any( + segment in (b".", b"..") for segment in unquote_to_bytes(path).split(b"/") + ) + + +def strip_access_token_from_uri(uri: bytes) -> bytes: + """Remove any `access_token` query parameter from the given request URI. + + Clients are not supposed to authenticate with an `access_token` query + parameter, but some might do so anyway. Since the URI is forwarded to the + application service verbatim, strip it out just in case, so that it isn't + inadvertently leaked to the application service. + """ + split_uri = urlsplit(uri) + if not split_uri.query: + return uri + + args = parse_qs(split_uri.query, keep_blank_values=True) + for key in list(args.keys()): + if key.lower() == b"access_token": + del args[key] + + if not args: + return split_uri.path + + return split_uri.path + b"?" + urlencode(args, doseq=True).encode("ascii") + + +async def proxy_request_to_appservice( + request: SynapseRequest, + hs: "HomeServer", + appservice: ApplicationService, + body_producer: Optional[IBodyProducer], + extra_request_headers: dict[bytes, bytes] | None = None, +) -> None: + """Forward the given request to an application service's proxy URL and stream + the response back to the original caller unchanged. + + Args: + request: The inbound request to forward. + hs: The homeserver. + appservice: The application service to forward the request to. Must have + `proxy_url` and `hs_token` set. + body_producer: A producer for the request body to forward, or None if the + request has no body to forward. + extra_request_headers: Additional headers to set on the outbound request, + beyond those copied from the original request. + """ + assert appservice.proxy_url is not None + assert appservice.hs_token is not None + + request_path = request.uri.split(b"?", 1)[0] + if has_dot_segments(request_path): + return_json_error( + failure.Failure( + SynapseError( + HTTPStatus.BAD_REQUEST, + "Request path is not allowed", + Codes.INVALID_PARAM, + ) + ), + request, + None, + ) + return + + target_uri = appservice.proxy_url.encode("ascii") + strip_access_token_from_uri( + request.uri + ) + + # Only forward the bare minimum of request headers an application service could + # plausibly need. + headers = Headers() + for header_name, header_values in request.requestHeaders.getAllRawHeaders(): + if header_name.decode("ascii").lower() in { + "content-type", + "accept", + "accept-language", + }: + headers.setRawHeaders(header_name, header_values) + + headers.setRawHeaders( + b"Authorization", [b"Bearer " + appservice.hs_token.encode("ascii")] + ) + + if extra_request_headers: + for header_name, header_value in extra_request_headers.items(): + headers.setRawHeaders(header_name, [header_value]) + + agent = hs.get_proxied_http_client().agent + request_deferred = run_in_background( + agent.request, + request.method, + target_uri, + headers=headers, + bodyProducer=body_producer, + ) + request_deferred = timeout_deferred( + deferred=request_deferred, + timeout=30, # Give the application service at most 30s to respond. + clock=hs.get_clock(), + ) + + try: + response = await make_deferred_yieldable(request_deferred) + except Exception: + logger.warning( + "Error proxying request to application service %s at %s", + appservice.id, + target_uri, + exc_info=True, + ) + return_json_error(failure.Failure(), request, None) + return + + _send_response(request, response) + + +def _send_response(request: SynapseRequest, response: IResponse) -> None: + response_headers = cast(Headers, response.headers) + + request.setResponseCode(response.code) + set_cors_headers(request) + + # We strip the "hop-by-hop" headers as defined by RFC2616. + headers_to_strip = set(HOP_BY_HOP_HEADERS_LOWERCASE) + + # The `Connection` header can define additional headers that should not be + # copied over. + connection_header = response_headers.getRawHeaders(b"connection") + headers_to_strip |= parse_connection_header_value( + connection_header[0] if connection_header else None + ) + + for header_name, header_values in response_headers.getAllRawHeaders(): + if header_name.decode("ascii").lower() in headers_to_strip: + continue + request.responseHeaders.setRawHeaders(header_name, header_values) + + response.deliverBody(_ProxyResponseBody(request)) diff --git a/synapse/http/proxy.py b/synapse/http/proxy.py index b3a2f84f29..ad5c6d5d91 100644 --- a/synapse/http/proxy.py +++ b/synapse/http/proxy.py @@ -51,16 +51,18 @@ logger = logging.getLogger(__name__) # "Hop-by-hop" headers (as opposed to "end-to-end" headers) as defined by RFC2616 # section 13.5.1 and referenced in RFC9110 section 7.6.1. These are meant to only be # consumed by the immediate recipient and not be forwarded on. -HOP_BY_HOP_HEADERS_LOWERCASE = { - "connection", - "keep-alive", - "proxy-authenticate", - "proxy-authorization", - "te", - "trailers", - "transfer-encoding", - "upgrade", -} +HOP_BY_HOP_HEADERS_LOWERCASE = frozenset( + { + "connection", + "keep-alive", + "proxy-authenticate", + "proxy-authorization", + "te", + "trailers", + "transfer-encoding", + "upgrade", + } +) assert all(header.lower() == header for header in HOP_BY_HOP_HEADERS_LOWERCASE) diff --git a/synapse/rest/__init__.py b/synapse/rest/__init__.py index c8ede662aa..1f705aee23 100644 --- a/synapse/rest/__init__.py +++ b/synapse/rest/__init__.py @@ -28,6 +28,7 @@ from synapse.rest.client import ( account_data, account_validity, appservice_ping, + appservice_proxy, auth, auth_metadata, capabilities, @@ -130,6 +131,7 @@ CLIENT_SERVLET_FUNCTIONS: tuple[RegisterServletsFunc, ...] = ( auth_metadata.register_servlets, thread_subscriptions.register_servlets, room_membership.register_servlets, + appservice_proxy.register_servlets, ) SERVLET_GROUPS: dict[str, Iterable[RegisterServletsFunc]] = { diff --git a/synapse/rest/client/appservice_proxy.py b/synapse/rest/client/appservice_proxy.py new file mode 100644 index 0000000000..aac1f465d3 --- /dev/null +++ b/synapse/rest/client/appservice_proxy.py @@ -0,0 +1,82 @@ +# +# 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 logging +import re +from typing import TYPE_CHECKING + +from synapse.api.ratelimiting import RequestRatelimiter +from synapse.appservice import ApplicationService +from synapse.http import QuieterFileBodyProducer +from synapse.http.appservice_proxy import proxy_request_to_appservice +from synapse.http.server import HttpServer, ServletCallback +from synapse.http.site import SynapseRequest + +if TYPE_CHECKING: + from synapse.server import HomeServer + +logger = logging.getLogger(__name__) + + +def _make_proxy_callback( + hs: "HomeServer", + ratelimiter: RequestRatelimiter, + appservice: ApplicationService, +) -> ServletCallback: + async def _proxy(request: SynapseRequest, **kwargs: str) -> None: + requester = await hs.get_auth().get_user_by_req(request) + + await ratelimiter.ratelimit(requester) + + await proxy_request_to_appservice( + request, + hs, + appservice, + QuieterFileBodyProducer(request.content), + extra_request_headers={ + b"X-Matrix-User-Identifier": requester.user.to_string().encode("ascii") + }, + ) + + return _proxy + + +def register_servlets(hs: "HomeServer", http_server: HttpServer) -> None: + """Registers blanket reverse-proxy routes for each application service that has + configured a proxy prefix. This forwards requests under + /_matrix/client///* (where is either "vN" or "unstable") + to the same path under the application service's proxy URL after verifying request + authentication. + """ + if not hs.config.experimental.msc4512_enabled: + return + + ratelimiter = hs.get_request_ratelimiter() + for appservice in hs.get_datastores().main.get_app_services(): + if appservice.proxy_prefix is None: + continue + + pattern = re.compile( + r"^/_matrix/client/(?:unstable/[^/]+|v[^/]+)/%s(/.*)?$" + % (re.escape(appservice.proxy_prefix),) + ) + callback = _make_proxy_callback(hs, ratelimiter, appservice) + + for method in ("GET", "POST", "PUT", "DELETE"): + http_server.register_paths( + method, + (pattern,), + callback, + "ApplicationServiceClientProxy", + ) diff --git a/tests/appservice/test_appservice.py b/tests/appservice/test_appservice.py index 3f124d9a2c..94bfdd3fba 100644 --- a/tests/appservice/test_appservice.py +++ b/tests/appservice/test_appservice.py @@ -291,3 +291,69 @@ class ApplicationServiceScopesTestCase(unittest.TestCase): token="some_token", scopes=["does:not:exist"], ) + + +class ApplicationServiceProxyPrefixTestCase(unittest.TestCase): + """Tests the proxying configuration for application services from MSC4512.""" + + def _make_service(self, **kwargs: Any) -> ApplicationService: + kwargs.setdefault("id", "unique_identifier") + kwargs.setdefault("sender", UserID.from_string("@as:test")) + kwargs.setdefault("token", "some_token") + return ApplicationService(**kwargs) + + def test_proxy_prefix_without_proxy_url_raises(self) -> None: + with self.assertRaises(KeyError): + self._make_service(proxy_url=None, proxy_prefix="rtc/livekit") + + def test_proxy_prefix_with_empty_proxy_url_raises(self) -> None: + with self.assertRaises(ValueError): + self._make_service(proxy_url="", proxy_prefix="rtc/livekit") + + def test_proxy_url_without_proxy_prefix_raises(self) -> None: + with self.assertRaises(KeyError): + self._make_service(proxy_url="http://proxy.example.com") + + def test_proxy_url_with_empty_proxy_prefix_raises(self) -> None: + with self.assertRaises(ValueError): + self._make_service(proxy_url="http://proxy.example.com", proxy_prefix="") + + def test_proxy_prefix_with_proxy_url_is_stored(self) -> None: + service = self._make_service( + proxy_url="http://proxy.example.com", + proxy_prefix="rtc/livekit", + ) + self.assertEqual(service.proxy_prefix, "rtc/livekit") + self.assertEqual(service.proxy_url, "http://proxy.example.com") + + def test_proxy_url_trailing_slash_is_stripped(self) -> None: + service = self._make_service( + proxy_url="http://proxy.example.com/", + proxy_prefix="rtc/livekit", + ) + self.assertEqual(service.proxy_url, "http://proxy.example.com") + + def test_nested_proxy_prefix_is_allowed(self) -> None: + service = self._make_service( + proxy_url="http://proxy.example.com", + proxy_prefix="rtc/livekit/foo", + ) + self.assertEqual(service.proxy_prefix, "rtc/livekit/foo") + + def test_disallowed_proxy_prefix_raises(self) -> None: + with self.assertRaises(ValueError): + self._make_service( + proxy_url="http://proxy.example.com", proxy_prefix="not/allowed" + ) + + def test_no_proxy_prefix_defaults_to_none(self) -> None: + service = self._make_service() + self.assertIsNone(service.proxy_prefix) + self.assertIsNone(service.proxy_url) + + def test_trailing_slash_on_proxy_prefix_is_stripped(self) -> None: + service = self._make_service( + proxy_url="http://proxy.example.com", + proxy_prefix="rtc/livekit/foo/", + ) + self.assertEqual(service.proxy_prefix, "rtc/livekit/foo") diff --git a/tests/federation/transport/server/test_appservice_proxy.py b/tests/federation/transport/server/test_appservice_proxy.py new file mode 100644 index 0000000000..957e1c3dae --- /dev/null +++ b/tests/federation/transport/server/test_appservice_proxy.py @@ -0,0 +1,349 @@ +# +# 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 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.server import HomeServer +from synapse.types import JsonDict +from synapse.util.clock import Clock + +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 ApplicationServiceFederationProxyTestCase(unittest.FederatingHomeserverTestCase): + """Tests the proxying of federation requests to application services from MSC4512.""" + + def default_config(self) -> JsonDict: + config = super().default_config() + _, path = tempfile.mkstemp(prefix="as_fed_proxy_config") + with open(path, "w") as f: + 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"] = [path] + config.setdefault("experimental_features", {}).setdefault( + "msc4512_enabled", True + ) + return config + + def prepare(self, reactor: MemoryReactor, clock: Clock, hs: HomeServer) -> None: + super().prepare(reactor, clock, hs) + self.agent = Mock() + hs.get_proxied_http_client().agent = self.agent + + def test_signed_get_is_proxied(self) -> None: + self.agent.request = Mock( + return_value=defer.succeed( + FakeResponse.json(code=200, payload={"ok": True}) + ) + ) + + channel = self.make_signed_federation_request( + "GET", f"/_matrix/federation/{VERSIONED_PREFIX}/some/path" + ) + + self.assertEqual(channel.code, 200) + self.assertEqual(channel.json_body, {"ok": True}) + + ((method, uri), kwargs) = self.agent.request.call_args + + self.assertEqual(method, b"GET") + self.assertEqual( + uri, + f"{APPSERVICE_URL}/_matrix/federation/{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-Origin"), + [self.OTHER_SERVER_NAME.encode("ascii")], + ) + + def test_signed_get_is_proxied_at_root_path(self) -> None: + self.agent.request = Mock( + return_value=defer.succeed( + FakeResponse.json(code=200, payload={"ok": True}) + ) + ) + + channel = self.make_signed_federation_request( + "GET", f"/_matrix/federation/{VERSIONED_PREFIX}" + ) + + self.assertEqual(channel.code, 200) + self.assertEqual(channel.json_body, {"ok": True}) + + ((method, uri), kwargs) = self.agent.request.call_args + + self.assertEqual(method, b"GET") + self.assertEqual( + uri, + f"{APPSERVICE_URL}/_matrix/federation/{VERSIONED_PREFIX}".encode(), + ) + + headers: Headers = kwargs["headers"] + self.assertEqual(headers.getRawHeaders(b"Authorization"), [b"Bearer hs_token"]) + self.assertEqual( + headers.getRawHeaders(b"X-Matrix-Origin"), + [self.OTHER_SERVER_NAME.encode("ascii")], + ) + + def test_signed_post_is_proxied(self) -> None: + self.agent.request = Mock( + return_value=defer.succeed( + FakeResponse.json(code=200, payload={"ok": True}) + ) + ) + + content = {"key": "value"} + channel = self.make_signed_federation_request( + "POST", + f"/_matrix/federation/{VERSIONED_PREFIX}/some/path", + content=content, + ) + + self.assertEqual(channel.code, 200) + self.assertEqual(channel.json_body, {"ok": True}) + + ((method, uri), kwargs) = self.agent.request.call_args + + self.assertEqual(method, b"POST") + self.assertEqual( + uri, + f"{APPSERVICE_URL}/_matrix/federation/{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-Origin"), + [self.OTHER_SERVER_NAME.encode("ascii")], + ) + self.assertEqual( + kwargs["headers"].getRawHeaders(b"Content-Type"), [b"application/json"] + ) + + body_producer = kwargs["bodyProducer"] + self.assertGreater(body_producer.length, 0) + + 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_signed_federation_request( + "GET", + f"/_matrix/federation/{VERSIONED_PREFIX}/some/path", + 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_signed_federation_request( + "GET", + f"/_matrix/federation/{VERSIONED_PREFIX}/some/path", + 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_signed_federation_request( + "POST", + f"/_matrix/federation/{VERSIONED_PREFIX}/some/path", + content={"key": "value"}, + 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_signed_federation_request( + "GET", f"/_matrix/federation/{VERSIONED_PREFIX}/some/path" + ) + + self.assertEqual(channel.code, 200) + self.assertEqual(channel.result["body"], b"hello") + self.assertEqual(channel.headers.getRawHeaders(b"X-Forward"), [b"forward"]) + + def test_unsigned_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/federation/{VERSIONED_PREFIX}/some/path", + shorthand=False, + ) + + self.assertEqual(channel.code, 401) + self.agent.request.assert_not_called() + + @unittest.override_config({"rc_federation": {"reject_limit": -1}}) + 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_signed_federation_request( + "GET", f"/_matrix/federation/{VERSIONED_PREFIX}/some/path" + ) + + self.assertEqual(channel.code, 429) + self.agent.request.assert_not_called() + + 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_signed_federation_request( + "GET", f"/_matrix/federation/{VERSIONED_PREFIX}/some/path" + ) + + self.assertEqual(channel.code, 500) + self.agent.request.assert_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_signed_federation_request( + "GET", f"/_matrix/federation/{VERSIONED_PREFIX}/some/../path" + ) + + 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_signed_federation_request( + "GET", f"/_matrix/federation/{VERSIONED_PREFIX}/some/%2e%2e/path" + ) + + 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_signed_federation_request( + "GET", "/_matrix/federation/not_a_registered_prefix/some/path" + ) + + self.assertEqual(channel.code, 404) + + def test_unregistered_prefix_with_suffix_is_rejected(self) -> None: + channel = self.make_signed_federation_request( + "GET", f"/_matrix/federation/{VERSIONED_PREFIX}-2" + ) + + self.assertEqual(channel.code, 404) + + def test_missing_version_segment_is_rejected(self) -> None: + channel = self.make_signed_federation_request( + "GET", f"/_matrix/federation/{APPSERVICE_PREFIX}/some/path" + ) + + self.assertEqual(channel.code, 404) + self.agent.request.assert_not_called() + + def test_unstable_version_segment_is_proxied(self) -> None: + self.agent.request = Mock( + return_value=defer.succeed( + FakeResponse.json(code=200, payload={"ok": True}) + ) + ) + + path = ( + f"/_matrix/federation/unstable/org.example.msc9999/" + f"{APPSERVICE_PREFIX}/some/path" + ) + channel = self.make_signed_federation_request("GET", path) + + self.assertEqual(channel.code, 200) + self.assertEqual(channel.json_body, {"ok": True}) + + ((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_signed_federation_request( + "GET", f"/_matrix/federation/{VERSIONED_PREFIX}/some/path" + ) + + self.assertEqual(channel.code, 404) + self.agent.request.assert_not_called() diff --git a/tests/http/test_appservice_proxy.py b/tests/http/test_appservice_proxy.py new file mode 100644 index 0000000000..de1ba3cae9 --- /dev/null +++ b/tests/http/test_appservice_proxy.py @@ -0,0 +1,47 @@ +# +# 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: +# . +# + +from synapse.http.appservice_proxy import has_dot_segments + +from tests import unittest + + +class HasDotSegmentsTestCase(unittest.TestCase): + def test_plain_path_has_no_dot_segments(self) -> None: + self.assertFalse(has_dot_segments(b"/some/path")) + self.assertFalse(has_dot_segments(b"/some/path.txt")) + self.assertFalse(has_dot_segments(b"/some/...path")) + + def test_dot_segment_is_detected(self) -> None: + self.assertTrue(has_dot_segments(b"/some/./path")) + self.assertTrue(has_dot_segments(b"/./some/path")) + self.assertTrue(has_dot_segments(b"/some/path/.")) + + def test_dot_dot_segment_is_detected(self) -> None: + self.assertTrue(has_dot_segments(b"/some/../path")) + self.assertTrue(has_dot_segments(b"/../some/path")) + self.assertTrue(has_dot_segments(b"/some/path/..")) + + def test_percent_encoded_dot_segments_are_detected(self) -> None: + self.assertTrue(has_dot_segments(b"/some/%2e%2e/path")) + self.assertTrue(has_dot_segments(b"/some/%2e/path")) + self.assertTrue(has_dot_segments(b"/some/%2E%2E/path")) + + def test_percent_encoded_separator_is_detected(self) -> None: + self.assertTrue(has_dot_segments(b"/some%2f../path")) + + def test_double_encoded_dot_segments_are_not_detected(self) -> None: + # Only a single decode is performed, matching the single decode that route + # arguments get elsewhere, so a double-encoded segment is left alone. + self.assertFalse(has_dot_segments(b"/some/%252e%252e/path")) diff --git a/tests/rest/client/test_appservice_proxy.py b/tests/rest/client/test_appservice_proxy.py new file mode 100644 index 0000000000..d389fa7d00 --- /dev/null +++ b/tests/rest/client/test_appservice_proxy.py @@ -0,0 +1,442 @@ +# +# 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() diff --git a/tests/storage/test_appservice.py b/tests/storage/test_appservice.py index 8430e54b92..1e097306e4 100644 --- a/tests/storage/test_appservice.py +++ b/tests/storage/test_appservice.py @@ -42,7 +42,7 @@ from synapse.storage.databases.main.appservice import ( ApplicationServiceStore, ApplicationServiceTransactionStore, ) -from synapse.types import DeviceListUpdates, JsonDict +from synapse.types import DeviceListUpdates from synapse.util.clock import Clock from tests import unittest @@ -483,8 +483,8 @@ class TestTransactionStore(ApplicationServiceTransactionStore, ApplicationServic class ApplicationServiceStoreConfigTestCase(unittest.HomeserverTestCase): - def _write_config(self, suffix: str, **kwargs: Any) -> str: - vals: JsonDict = { + def _write_config(self, suffix: str, **kwargs: str | list[str] | None) -> str: + vals: dict[str, Any] = { "id": "id" + suffix, "url": "url" + suffix, "as_token": "as_token" + suffix, @@ -636,3 +636,231 @@ class ApplicationServiceStoreConfigTestCase(unittest.HomeserverTestCase): ), self.hs, ) + + def test_proxy_prefix_works(self) -> None: + f1 = self._write_config( + suffix="1", + **{ + "io.element.msc4512.proxy_prefix": "rtc/livekit", + "io.element.msc4512.proxy_url": "http://proxy", + }, + ) + + self.hs.config.appservice.app_service_config_files = [f1] + self.hs.config.caches.event_cache_size = 1 + + server_name = self.hs.hostname + database = self.hs.get_datastores().databases[0] + store = ApplicationServiceStore( + database, + make_conn( + db_config=database._database_config, + engine=database.engine, + default_txn_name="test", + server_name=server_name, + ), + self.hs, + ) + (appservice,) = store.get_app_services() + self.assertEqual(appservice.proxy_prefix, "rtc/livekit") + self.assertEqual(appservice.proxy_url, "http://proxy") + + def test_proxy_prefix_requires_proxy_url(self) -> None: + f1 = self._write_config( + suffix="1", + **{"io.element.msc4512.proxy_prefix": "rtc/livekit"}, + ) + + self.hs.config.appservice.app_service_config_files = [f1] + self.hs.config.caches.event_cache_size = 1 + + with self.assertRaises(KeyError): + server_name = self.hs.hostname + database = self.hs.get_datastores().databases[0] + ApplicationServiceStore( + database, + make_conn( + db_config=database._database_config, + engine=database.engine, + default_txn_name="test", + server_name=server_name, + ), + self.hs, + ) + + def test_proxy_url_requires_proxy_prefix(self) -> None: + f1 = self._write_config( + suffix="1", + **{"io.element.msc4512.proxy_url": "http://proxy"}, + ) + + self.hs.config.appservice.app_service_config_files = [f1] + self.hs.config.caches.event_cache_size = 1 + + with self.assertRaises(KeyError): + server_name = self.hs.hostname + database = self.hs.get_datastores().databases[0] + ApplicationServiceStore( + database, + make_conn( + db_config=database._database_config, + engine=database.engine, + default_txn_name="test", + server_name=server_name, + ), + self.hs, + ) + + def test_proxy_prefix_requires_non_empty_proxy_url(self) -> None: + f1 = self._write_config( + suffix="1", + **{ + "io.element.msc4512.proxy_prefix": "rtc/livekit", + "io.element.msc4512.proxy_url": "", + }, + ) + + self.hs.config.appservice.app_service_config_files = [f1] + self.hs.config.caches.event_cache_size = 1 + + with self.assertRaises(ValueError): + server_name = self.hs.hostname + database = self.hs.get_datastores().databases[0] + ApplicationServiceStore( + database, + make_conn( + db_config=database._database_config, + engine=database.engine, + default_txn_name="test", + server_name=server_name, + ), + self.hs, + ) + + def test_proxy_url_requires_non_empty_proxy_prefix(self) -> None: + f1 = self._write_config( + suffix="1", + **{ + "io.element.msc4512.proxy_prefix": "", + "io.element.msc4512.proxy_url": "http://proxy", + }, + ) + + self.hs.config.appservice.app_service_config_files = [f1] + self.hs.config.caches.event_cache_size = 1 + + with self.assertRaises(ValueError): + server_name = self.hs.hostname + database = self.hs.get_datastores().databases[0] + ApplicationServiceStore( + database, + make_conn( + db_config=database._database_config, + engine=database.engine, + default_txn_name="test", + server_name=server_name, + ), + self.hs, + ) + + def test_proxy_prefix_does_not_allow_reserved_values(self) -> None: + f1 = self._write_config( + suffix="1", + **{ + "io.element.msc4512.proxy_prefix": "not/allowed", + "io.element.msc4512.proxy_url": "http://proxy", + }, + ) + + self.hs.config.appservice.app_service_config_files = [f1] + self.hs.config.caches.event_cache_size = 1 + + with self.assertRaises(ValueError): + server_name = self.hs.hostname + database = self.hs.get_datastores().databases[0] + ApplicationServiceStore( + database, + make_conn( + db_config=database._database_config, + engine=database.engine, + default_txn_name="test", + server_name=server_name, + ), + self.hs, + ) + + def test_duplicate_proxy_prefix(self) -> None: + f1 = self._write_config( + suffix="1", + **{ + "io.element.msc4512.proxy_prefix": "rtc/livekit", + "io.element.msc4512.proxy_url": "http://proxy", + }, + ) + f2 = self._write_config( + suffix="2", + **{ + "io.element.msc4512.proxy_prefix": "rtc/livekit", + "io.element.msc4512.proxy_url": "http://proxy2", + }, + ) + + self.hs.config.appservice.app_service_config_files = [f1, f2] + self.hs.config.caches.event_cache_size = 1 + + with self.assertRaises(ConfigError) as cm: + server_name = self.hs.hostname + database = self.hs.get_datastores().databases[0] + ApplicationServiceStore( + database, + make_conn( + db_config=database._database_config, + engine=database.engine, + default_txn_name="test", + server_name=server_name, + ), + self.hs, + ) + + e = cm.exception + self.assertIn(f1, str(e)) + self.assertIn(f2, str(e)) + self.assertIn("io.element.msc4512.proxy_prefix", str(e)) + + def test_overlapping_proxy_prefix(self) -> None: + f1 = self._write_config( + suffix="1", + **{ + "io.element.msc4512.proxy_prefix": "rtc/livekit", + "io.element.msc4512.proxy_url": "http://proxy", + }, + ) + f2 = self._write_config( + suffix="2", + **{ + "io.element.msc4512.proxy_prefix": "rtc/livekit/foobar", + "io.element.msc4512.proxy_url": "http://proxy2", + }, + ) + + self.hs.config.appservice.app_service_config_files = [f1, f2] + self.hs.config.caches.event_cache_size = 1 + + with self.assertRaises(ConfigError) as cm: + server_name = self.hs.hostname + database = self.hs.get_datastores().databases[0] + ApplicationServiceStore( + database, + make_conn( + db_config=database._database_config, + engine=database.engine, + default_txn_name="test", + server_name=server_name, + ), + self.hs, + ) + + e = cm.exception + self.assertIn(f1, str(e)) + self.assertIn(f2, str(e)) + self.assertIn("io.element.msc4512.proxy_prefix", str(e))