diff --git a/changelog.d/19972.feature b/changelog.d/19972.feature new file mode 100644 index 00000000000..56e3b271368 --- /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/changelog.d/19977.feature b/changelog.d/19977.feature new file mode 100644 index 00000000000..cd18e36f69c --- /dev/null +++ b/changelog.d/19977.feature @@ -0,0 +1 @@ +Add experimental support for sending federation requests from app services as per MSC4512. diff --git a/synapse/api/errors.py b/synapse/api/errors.py index 0c35b4a7ba5..18ac13bf23a 100644 --- a/synapse/api/errors.py +++ b/synapse/api/errors.py @@ -125,6 +125,11 @@ class Codes(str, Enum): AS_PING_CONNECTION_TIMEOUT = "M_CONNECTION_TIMEOUT" AS_PING_CONNECTION_FAILED = "M_CONNECTION_FAILED" + AS_FEDPROXY_NO_PROXY_PREFIX = "IO.ELEMENT.MSC4512.M_FEDPROXY_NO_PROXY_PREFIX" + AS_FEDPROXY_PATH_NOT_ALLOWED = "IO.ELEMENT.MSC4512.M_FEDPROXY_PATH_NOT_ALLOWED" + AS_FEDPROXY_CONNECTION_FAILED = "IO.ELEMENT.MSC4512.M_FEDPROXY_CONNECTION_FAILED" + AS_FEDPROXY_DESTINATION_DENIED = "IO.ELEMENT.MSC4512.M_FEDPROXY_DESTINATION_DENIED" + # Attempt to send a second annotation with the same event type & annotation key # MSC2677 DUPLICATE_ANNOTATION = "M_DUPLICATE_ANNOTATION" diff --git a/synapse/appservice/__init__.py b/synapse/appservice/__init__.py index c55a83a8799..611efaa5651 100644 --- a/synapse/appservice/__init__.py +++ b/synapse/appservice/__init__.py @@ -89,6 +89,8 @@ class ApplicationService: # values. NS_LIST = [NS_USERS, NS_ALIASES, NS_ROOMS] + ALLOWED_PROXY_PREFIXES = {"unstable/io.element.msc4195/rtc/livekit"} + def __init__( self, token: str, @@ -104,6 +106,7 @@ def __init__( supports_unstable_ephemeral: bool = False, msc3202_transaction_extensions: bool = False, msc4190_device_management: bool = False, + proxy_prefix: str | None = None, ): self.token = token self.url = ( @@ -130,10 +133,17 @@ def __init__( self.supports_ephemeral = supports_ephemeral self.msc3202_transaction_extensions = msc3202_transaction_extensions self.msc4190_device_management = msc4190_device_management + self.proxy_prefix = proxy_prefix if "|" in self.id: raise Exception("application service ID cannot contain '|' character") + if proxy_prefix is not None: + if not self._is_proxy_prefix_allowed(proxy_prefix): + raise ValueError(f"cannot claim reserved proxy prefix {proxy_prefix}") + if not self.url: + raise KeyError("cannot claim proxy prefix without also setting a url") + # .protocols is a publicly visible field if protocols: self.protocols = set(protocols) @@ -192,6 +202,12 @@ def _is_exclusive(self, namespace_key: str, test_string: str) -> bool: 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 7a629d10bf6..767467d80d9 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,18 @@ def load_appservices( ) ) seen_as_tokens[appservice.token] = config_file + if appservice.proxy_prefix is not None: + if appservice.proxy_prefix in seen_proxy_prefixes: + raise ConfigError( + "Cannot reuse io.element.msc4512.proxy across application services: " + "%s (files: %s, %s)" + % ( + appservice.proxy_prefix, + config_file, + seen_proxy_prefixes[appservice.proxy_prefix], + ) + ) + seen_proxy_prefixes[appservice.proxy_prefix] = config_file logger.info("Loaded application service: %s", appservice) appservices.append(appservice) except Exception as e: @@ -199,6 +212,17 @@ def _load_appservice( "The `io.element.msc4190` option should be true or false 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: + # - /_matrix/client/{prefix}/* -> {url}/_matrix/client/{prefix}/* + # - /_matrix/federation/{prefix}/* -> {url}/_matrix/federation//* + proxy_prefix = as_info.get("io.element.msc4512.proxy") + if proxy_prefix is not None: + if not isinstance(proxy_prefix, str) or not proxy_prefix: + raise ValueError( + "The `io.element.msc4512.proxy` option should be a non-empty string." + ) + return ApplicationService( token=as_info["as_token"], url=as_info["url"], @@ -213,4 +237,5 @@ def _load_appservice( supports_ephemeral=supports_ephemeral, msc3202_transaction_extensions=msc3202_transaction_extensions, msc4190_device_management=msc4190_enabled, + proxy_prefix=proxy_prefix, ) diff --git a/synapse/config/experimental.py b/synapse/config/experimental.py index 6ad9f535179..794b5f0f419 100644 --- a/synapse/config/experimental.py +++ b/synapse/config/experimental.py @@ -309,3 +309,6 @@ def read_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 0eff49cf73c..70fc55123ba 100644 --- a/synapse/federation/transport/server/__init__.py +++ b/synapse/federation/transport/server/__init__.py @@ -24,6 +24,7 @@ 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 00000000000..242dec73500 --- /dev/null +++ b/synapse/federation/transport/server/appservice_proxy.py @@ -0,0 +1,106 @@ +# +# 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: + async def _proxy(request: SynapseRequest, **kwargs: str) -> None: + raw_body = request.content.read() # type: ignore[union-attr] + + content = None + if request.method in (b"PUT", b"POST"): + try: + content = json_decoder.decode(raw_body.decode("utf-8")) + except Exception: + raise SynapseError( + HTTPStatus.BAD_REQUEST, "Content not JSON.", Codes.NOT_JSON + ) + + 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, + QuieterFileBodyProducer(BytesIO(raw_body)), + 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//* + to the same path under the application service's 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: + continue + + pattern = re.compile( + "^/_matrix/federation/%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 00000000000..1beaa18bfe1 --- /dev/null +++ b/synapse/http/appservice_proxy.py @@ -0,0 +1,248 @@ +# +# 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 json +import logging +from http import HTTPStatus +from typing import TYPE_CHECKING, cast + +from twisted.web.http_headers import Headers +from twisted.web.iweb import IBodyProducer, IResponse + +from synapse.api.errors import ( + Codes, + FederationDeniedError, + HttpResponseException, + RequestSendFailed, + 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 set_cors_headers +from synapse.http.site import SynapseRequest +from synapse.http.types import QueryParams +from synapse.logging.context import make_deferred_yieldable, run_in_background +from synapse.types import JsonDict +from synapse.util.async_helpers import timeout_deferred +from synapse.util.json import json_decoder +from synapse.util.retryutils import NotRetryingDestination + +if TYPE_CHECKING: + from synapse.server import HomeServer + +logger = logging.getLogger(__name__) + + +async def proxy_request_to_appservice( + request: SynapseRequest, + hs: "HomeServer", + appservice: ApplicationService, + body_producer: IBodyProducer, + extra_request_headers: dict[bytes, bytes] | None = None, +) -> None: + """Forward the given request to an application service's 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 + `url` set. + body_producer: A producer for the request body to forward. + extra_request_headers: Additional headers to set on the outbound request, + beyond those copied from the original request. + """ + assert appservice.url is not None + target_uri = appservice.url.encode("ascii") + request.uri + + # Other than the "hop-by-hop" headers as defined by RFC2616 we also strip: + # - Host and Content-Length (twisted adds these on its own) + # - Authorization (because the app service shouldn't need to be concerned with it) + headers_to_strip = HOP_BY_HOP_HEADERS_LOWERCASE | { + "host", + "content-length", + "authorization", + } + + # The `Connection` header can define additional headers that should not be + # copied over. + connection_header = request.requestHeaders.getRawHeaders(b"connection") + headers_to_strip |= parse_connection_header_value( + connection_header[0] if connection_header else None + ) + + headers = Headers() + for header_name, header_values in request.requestHeaders.getAllRawHeaders(): + if header_name.decode("ascii").lower() in headers_to_strip: + continue + headers.setRawHeaders(header_name, header_values) + 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, + ) + _send_error_response(request) + 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 = 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)) + + +def _send_error_response(request: SynapseRequest) -> None: + request.setResponseCode(404) + set_cors_headers(request) + request.setHeader(b"Content-Type", b"application/json") + request.write( + json.dumps( + {"errcode": Codes.UNRECOGNIZED, "error": "Unrecognized request"} + ).encode() + ) + request.finish() + + +async def send_federation_request_from_appservice( + hs: "HomeServer", + appservice: ApplicationService, + method: str, + destination: str, + path: str, + data: JsonDict | None, + args: QueryParams | None, +) -> tuple[int, JsonDict | None]: + """Sign and send a federation request on behalf of an appservice. + + Returns: + A `(status, content)` tuple describing the destination's actual HTTP response. + """ + _check_path_allowed_for_appservice(appservice, path) + + if destination == hs.hostname: + raise SynapseError( + HTTPStatus.FORBIDDEN, + "Cannot target this homeserver itself", + Codes.AS_FEDPROXY_DESTINATION_DENIED, + ) + + client = hs.get_federation_http_client() + + try: + if method == "GET": + content = await client.get_json(destination, path, args=args) + elif method == "PUT": + content = await client.put_json(destination, path, args=args, data=data) + elif method == "POST": + content = await client.post_json(destination, path, args=args, data=data) + elif method == "DELETE": + content = await client.delete_json(destination, path, args=args) + else: + raise SynapseError( + HTTPStatus.BAD_REQUEST, + f"Unsupported method {method}", + Codes.INVALID_PARAM, + ) + return HTTPStatus.OK, content + except HttpResponseException as e: + try: + content = json_decoder.decode(e.response.decode("utf-8")) + except (UnicodeDecodeError, ValueError): + content = None + return e.code, content + except FederationDeniedError as e: + raise SynapseError( + HTTPStatus.FORBIDDEN, + e.msg, + Codes.AS_FEDPROXY_DESTINATION_DENIED, + ) + except (RequestSendFailed, NotRetryingDestination) as e: + raise SynapseError( + HTTPStatus.BAD_GATEWAY, + str(e), + Codes.AS_FEDPROXY_CONNECTION_FAILED, + ) + + +def _check_path_allowed_for_appservice( + appservice: ApplicationService, path: str +) -> None: + # Deny relative paths. + if any(segment in (".", "..") for segment in path.split("/")): + raise SynapseError( + HTTPStatus.FORBIDDEN, + "Path must not contain '.' or '..' segments", + Codes.AS_FEDPROXY_PATH_NOT_ALLOWED, + ) + + # Ensure the path is under the appservice's own proxy prefix. + if appservice.proxy_prefix is None: + raise SynapseError( + HTTPStatus.BAD_REQUEST, + "Application service does not have a proxy prefix", + Codes.AS_FEDPROXY_NO_PROXY_PREFIX, + ) + allowed_root = f"/_matrix/federation/{appservice.proxy_prefix}" + if path != allowed_root and not path.startswith(allowed_root + "/"): + raise SynapseError( + HTTPStatus.FORBIDDEN, + f"Path must be under {allowed_root}", + Codes.AS_FEDPROXY_PATH_NOT_ALLOWED, + ) diff --git a/synapse/rest/__init__.py b/synapse/rest/__init__.py index a56a81a8e91..90489a02ee7 100644 --- a/synapse/rest/__init__.py +++ b/synapse/rest/__init__.py @@ -27,7 +27,9 @@ account, account_data, account_validity, + appservice_federation_proxy, appservice_ping, + appservice_proxy, auth, auth_metadata, capabilities, @@ -128,6 +130,8 @@ rendezvous.register_servlets, auth_metadata.register_servlets, thread_subscriptions.register_servlets, + appservice_proxy.register_servlets, + appservice_federation_proxy.register_servlets, ) SERVLET_GROUPS: dict[str, Iterable[RegisterServletsFunc]] = { diff --git a/synapse/rest/client/appservice_federation_proxy.py b/synapse/rest/client/appservice_federation_proxy.py new file mode 100644 index 00000000000..379b23740d4 --- /dev/null +++ b/synapse/rest/client/appservice_federation_proxy.py @@ -0,0 +1,137 @@ +# +# 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 typing import TYPE_CHECKING + +from synapse.api.errors import Codes, SynapseError +from synapse.http.appservice_proxy import send_federation_request_from_appservice +from synapse.http.server import HttpServer +from synapse.http.servlet import RestServlet, parse_json_object_from_request +from synapse.http.site import SynapseRequest +from synapse.types import JsonDict + +if TYPE_CHECKING: + from synapse.server import HomeServer + +logger = logging.getLogger(__name__) + +ALLOWED_METHODS = ("GET", "PUT", "POST", "DELETE") + + +class AppserviceFederationProxyRestServlet(RestServlet): + PATTERNS = [ + re.compile( + r"^/_matrix/client/unstable/io.element.msc4512/appservice/fed_proxy$" + ) + ] + + def __init__(self, hs: "HomeServer"): + super().__init__() + self.hs = hs + self.auth = hs.get_auth() + self.store = hs.get_datastores().main + + async def on_POST(self, request: SynapseRequest) -> tuple[int, JsonDict]: + requester = await self.auth.get_user_by_req(request) + + app_service = ( + self.store.get_app_service_by_id(requester.app_service_id) + if requester.app_service_id + else None + ) + + if not app_service: + raise SynapseError( + HTTPStatus.FORBIDDEN, + "Only application services can use this endpoint", + Codes.FORBIDDEN, + ) + + content = parse_json_object_from_request(request) + + destination = content.get("destination") + if not isinstance(destination, str) or not destination: + raise SynapseError( + HTTPStatus.BAD_REQUEST, + "Missing or invalid destination", + Codes.MISSING_PARAM, + ) + + method = content.get("method") + if method not in ALLOWED_METHODS: + raise SynapseError( + HTTPStatus.BAD_REQUEST, + f"method must be one of {ALLOWED_METHODS}", + Codes.INVALID_PARAM, + ) + + path = content.get("path") + if not isinstance(path, str) or not path: + raise SynapseError( + HTTPStatus.BAD_REQUEST, + "Missing or invalid path", + Codes.MISSING_PARAM, + ) + + body = content.get("body") + if body is not None and not isinstance(body, dict): + raise SynapseError( + HTTPStatus.BAD_REQUEST, + "body must be an object", + Codes.INVALID_PARAM, + ) + if body is not None and method in ("GET", "DELETE"): + raise SynapseError( + HTTPStatus.BAD_REQUEST, + f"'body' is not supported for {method}", + Codes.INVALID_PARAM, + ) + + query = content.get("query") + if query is not None and not isinstance(query, dict): + raise SynapseError( + HTTPStatus.BAD_REQUEST, + "query must be an object", + Codes.INVALID_PARAM, + ) + if query is not None and any( + not isinstance(value, str) for value in query.values() + ): + raise SynapseError( + HTTPStatus.BAD_REQUEST, + "'query' values must be strings", + Codes.INVALID_PARAM, + ) + + status, response_content = await send_federation_request_from_appservice( + self.hs, + app_service, + method, + destination, + path, + body, + query, + ) + + return HTTPStatus.OK, {"status": status, "content": response_content} + + +def register_servlets(hs: "HomeServer", http_server: HttpServer) -> None: + if not hs.config.experimental.msc4512_enabled: + return + + AppserviceFederationProxyRestServlet(hs).register(http_server) diff --git a/synapse/rest/client/appservice_proxy.py b/synapse/rest/client/appservice_proxy.py new file mode 100644 index 00000000000..63f84efa059 --- /dev/null +++ b/synapse/rest/client/appservice_proxy.py @@ -0,0 +1,80 @@ +# +# 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//* + to the same path under the application service's 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( + "^/_matrix/client/%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 620c2b907b2..227c513ba1f 100644 --- a/tests/appservice/test_appservice.py +++ b/tests/appservice/test_appservice.py @@ -257,3 +257,49 @@ def test_member_list_match(self) -> Generator["defer.Deferred[Any]", object, Non ) ) ) + + +class ApplicationServiceProxyPrefixTestCase(unittest.TestCase): + 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_url_raises(self) -> None: + with self.assertRaises(KeyError): + self._make_service( + url=None, proxy_prefix="unstable/io.element.msc4195/rtc/livekit" + ) + + def test_proxy_prefix_with_empty_url_raises(self) -> None: + with self.assertRaises(KeyError): + self._make_service( + url="", proxy_prefix="unstable/io.element.msc4195/rtc/livekit" + ) + + def test_proxy_prefix_with_url_is_stored(self) -> None: + service = self._make_service( + url="http://example.com", + proxy_prefix="unstable/io.element.msc4195/rtc/livekit", + ) + self.assertEqual( + service.proxy_prefix, "unstable/io.element.msc4195/rtc/livekit" + ) + + def test_nested_proxy_prefix_is_allowed(self) -> None: + service = self._make_service( + url="http://example.com", + proxy_prefix="unstable/io.element.msc4195/rtc/livekit/foo", + ) + self.assertEqual( + service.proxy_prefix, "unstable/io.element.msc4195/rtc/livekit/foo" + ) + + def test_disallowed_proxy_prefix_raises(self) -> None: + with self.assertRaises(ValueError): + self._make_service(url="http://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) 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 00000000000..1f5374934fe --- /dev/null +++ b/tests/federation/transport/server/test_appservice_proxy.py @@ -0,0 +1,290 @@ +# +# 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 = "unstable/io.element.msc4195/rtc/livekit" + + +class ApplicationServiceFederationProxyTestCase(unittest.FederatingHomeserverTestCase): + 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": APPSERVICE_URL, + "as_token": "as_token", + "hs_token": "hs_token", + "sender_localpart": "proxy_bot", + "namespaces": {}, + "io.element.msc4512.proxy": APPSERVICE_PREFIX, + }, + 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/{APPSERVICE_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/{APPSERVICE_PREFIX}/some/path".encode(), + ) + + headers: Headers = kwargs["headers"] + self.assertIsNone(headers.getRawHeaders(b"Authorization")) + 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/{APPSERVICE_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/{APPSERVICE_PREFIX}".encode(), + ) + + headers: Headers = kwargs["headers"] + self.assertIsNone(headers.getRawHeaders(b"Authorization")) + 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/{APPSERVICE_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/{APPSERVICE_PREFIX}/some/path".encode(), + ) + + headers: Headers = kwargs["headers"] + self.assertIsNone(headers.getRawHeaders(b"Authorization")) + self.assertEqual( + headers.getRawHeaders(b"X-Matrix-Origin"), + [self.OTHER_SERVER_NAME.encode("ascii")], + ) + + body_producer = kwargs["bodyProducer"] + self.assertGreater(body_producer.length, 0) + + def test_connection_header_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/{APPSERVICE_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.assertEqual(headers.getRawHeaders(b"X-Forward"), [b"forward"]) + + def test_connection_header_with_named_request_headers_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/{APPSERVICE_PREFIX}/some/path", + custom_headers=[ + ("Connection", "close, X-Omit"), + ("X-Omit", "omit"), + ("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-Omit")) + self.assertEqual(headers.getRawHeaders(b"X-Forward"), [b"forward"]) + + 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/{APPSERVICE_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/{APPSERVICE_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/{APPSERVICE_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/{APPSERVICE_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/{APPSERVICE_PREFIX}/some/path" + ) + + self.assertEqual(channel.code, 404) + self.agent.request.assert_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/{APPSERVICE_PREFIX}-2" + ) + + self.assertEqual(channel.code, 404) + + @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/{APPSERVICE_PREFIX}/some/path" + ) + + self.assertEqual(channel.code, 404) + self.agent.request.assert_not_called() diff --git a/tests/rest/client/test_appservice_federation_proxy.py b/tests/rest/client/test_appservice_federation_proxy.py new file mode 100644 index 00000000000..c28b83a50ab --- /dev/null +++ b/tests/rest/client/test_appservice_federation_proxy.py @@ -0,0 +1,435 @@ +# +# 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 json +from unittest import mock + +from twisted.internet.testing import MemoryReactor + +from synapse.api.errors import Codes +from synapse.appservice import ApplicationService +from synapse.rest import admin +from synapse.rest.client import appservice_federation_proxy, login +from synapse.server import HomeServer +from synapse.types import JsonDict, UserID +from synapse.util.clock import Clock + +from tests import unittest +from tests.server import FakeChannel +from tests.test_utils import FakeResponse + +APPSERVICE_URL = "http://appservice.example.com" +APPSERVICE_PREFIX = "unstable/io.element.msc4195/rtc/livekit" +AS_TOKEN = "as_token" + + +class ApplicationServiceFederationProxyTestCase(unittest.HomeserverTestCase): + servlets = [ + admin.register_servlets, + login.register_servlets, + appservice_federation_proxy.register_servlets, + ] + + def default_config(self) -> JsonDict: + config = super().default_config() + config.setdefault("experimental_features", {}).setdefault( + "msc4512_enabled", True + ) + return config + + def prepare(self, reactor: MemoryReactor, clock: Clock, hs: HomeServer) -> None: + self.appservice = ApplicationService( + AS_TOKEN, + id="proxy_as", + sender=UserID.from_string("@proxy_bot:test"), + namespaces={}, + url=APPSERVICE_URL, + proxy_prefix=APPSERVICE_PREFIX, + ) + hs.get_datastores().main.services_cache.append(self.appservice) + + self.agent_request = mock.AsyncMock() + hs.get_federation_http_client().agent.request = self.agent_request # type: ignore[method-assign] + + def _fed_proxy( + self, content: dict, access_token: str | None = AS_TOKEN + ) -> FakeChannel: + return self.make_request( + "POST", + "/_matrix/client/unstable/io.element.msc4512/appservice/fed_proxy", + content, + access_token=access_token, + ) + + def test_get_is_sent_and_relayed(self) -> None: + self.agent_request.return_value = FakeResponse.json( + code=200, payload={"hello": "world"} + ) + + channel = self._fed_proxy( + { + "destination": "remote.example.com", + "method": "GET", + "path": f"/_matrix/federation/{APPSERVICE_PREFIX}/some/path", + "query": {"foo": "bar"}, + } + ) + + self.assertEqual(channel.code, 200) + self.assertEqual( + channel.json_body, {"status": 200, "content": {"hello": "world"}} + ) + + ((method, uri), kwargs) = self.agent_request.call_args + self.assertEqual(method, b"GET") + self.assertEqual( + uri, + f"matrix-federation://remote.example.com/_matrix/federation/{APPSERVICE_PREFIX}/some/path?foo=bar".encode(), + ) + + self.assertIsNone(kwargs["bodyProducer"]) + + headers = kwargs["headers"] + expected_auth_headers = self.hs.get_federation_http_client().build_auth_headers( + b"remote.example.com", + b"GET", + f"/_matrix/federation/{APPSERVICE_PREFIX}/some/path?foo=bar".encode(), + ) + self.assertEqual(headers.getRawHeaders(b"Authorization"), expected_auth_headers) + + def test_delete_is_sent_and_relayed(self) -> None: + self.agent_request.return_value = FakeResponse.json( + code=200, payload={"hello": "world"} + ) + + channel = self._fed_proxy( + { + "destination": "remote.example.com", + "method": "DELETE", + "path": f"/_matrix/federation/{APPSERVICE_PREFIX}/some/path", + "query": {"foo": "bar"}, + } + ) + + self.assertEqual(channel.code, 200) + self.assertEqual( + channel.json_body, {"status": 200, "content": {"hello": "world"}} + ) + + ((method, uri), kwargs) = self.agent_request.call_args + self.assertEqual(method, b"DELETE") + self.assertEqual( + uri, + f"matrix-federation://remote.example.com/_matrix/federation/{APPSERVICE_PREFIX}/some/path?foo=bar".encode(), + ) + + self.assertIsNone(kwargs["bodyProducer"]) + + headers = kwargs["headers"] + expected_auth_headers = self.hs.get_federation_http_client().build_auth_headers( + b"remote.example.com", + b"DELETE", + f"/_matrix/federation/{APPSERVICE_PREFIX}/some/path?foo=bar".encode(), + ) + self.assertEqual(headers.getRawHeaders(b"Authorization"), expected_auth_headers) + + def test_post_with_body_is_sent_and_relayed(self) -> None: + self.agent_request.return_value = FakeResponse.json( + code=200, payload={"hello": "world"} + ) + + channel = self._fed_proxy( + { + "destination": "remote.example.com", + "method": "POST", + "path": f"/_matrix/federation/{APPSERVICE_PREFIX}/some/path", + "body": {"key": "value"}, + "query": {"foo": "bar"}, + } + ) + + self.assertEqual(channel.code, 200) + self.assertEqual( + channel.json_body, {"status": 200, "content": {"hello": "world"}} + ) + + ((method, uri), kwargs) = self.agent_request.call_args + self.assertEqual(method, b"POST") + self.assertEqual( + uri, + f"matrix-federation://remote.example.com/_matrix/federation/{APPSERVICE_PREFIX}/some/path?foo=bar".encode(), + ) + + body_producer = kwargs["bodyProducer"] + self.assertEqual( + json.loads(body_producer._inputFile.getvalue()), + {"key": "value"}, + ) + + headers = kwargs["headers"] + expected_auth_headers = self.hs.get_federation_http_client().build_auth_headers( + b"remote.example.com", + b"POST", + f"/_matrix/federation/{APPSERVICE_PREFIX}/some/path?foo=bar".encode(), + content={"key": "value"}, + ) + self.assertEqual(headers.getRawHeaders(b"Authorization"), expected_auth_headers) + + def test_put_with_body_is_sent_and_relayed(self) -> None: + self.agent_request.return_value = FakeResponse.json( + code=200, payload={"hello": "world"} + ) + + channel = self._fed_proxy( + { + "destination": "remote.example.com", + "method": "PUT", + "path": f"/_matrix/federation/{APPSERVICE_PREFIX}/some/path", + "body": {"key": "value"}, + "query": {"foo": "bar"}, + } + ) + + self.assertEqual(channel.code, 200) + self.assertEqual( + channel.json_body, {"status": 200, "content": {"hello": "world"}} + ) + + ((method, uri), kwargs) = self.agent_request.call_args + self.assertEqual(method, b"PUT") + self.assertEqual( + uri, + f"matrix-federation://remote.example.com/_matrix/federation/{APPSERVICE_PREFIX}/some/path?foo=bar".encode(), + ) + + body_producer = kwargs["bodyProducer"] + self.assertEqual( + json.loads(body_producer._inputFile.getvalue()), + {"key": "value"}, + ) + + headers = kwargs["headers"] + expected_auth_headers = self.hs.get_federation_http_client().build_auth_headers( + b"remote.example.com", + b"PUT", + f"/_matrix/federation/{APPSERVICE_PREFIX}/some/path?foo=bar".encode(), + content={"key": "value"}, + ) + self.assertEqual(headers.getRawHeaders(b"Authorization"), expected_auth_headers) + + def test_remote_error_response_is_relayed(self) -> None: + self.agent_request.return_value = FakeResponse.json( + code=400, payload={"errcode": "M_UNRECOGNIZED"} + ) + + channel = self._fed_proxy( + { + "destination": "remote.example.com", + "method": "GET", + "path": f"/_matrix/federation/{APPSERVICE_PREFIX}/some/path", + } + ) + + self.assertEqual(channel.code, 200) + self.assertEqual( + channel.json_body, + {"status": 400, "content": {"errcode": "M_UNRECOGNIZED"}}, + ) + + @unittest.override_config({"federation": {"max_short_retries": 0}}) + def test_connection_failure_causes_502(self) -> None: + self.agent_request.side_effect = Exception("boom") + + channel = self._fed_proxy( + { + "destination": "remote.example.com", + "method": "GET", + "path": f"/_matrix/federation/{APPSERVICE_PREFIX}/some/path", + } + ) + + self.assertEqual(channel.code, 502) + self.assertEqual( + channel.json_body["errcode"], Codes.AS_FEDPROXY_CONNECTION_FAILED + ) + + def test_denied_destination_is_rejected(self) -> None: + channel = self._fed_proxy( + { + "destination": "not a valid server name", + "method": "GET", + "path": f"/_matrix/federation/{APPSERVICE_PREFIX}/some/path", + } + ) + + self.assertEqual(channel.code, 403) + self.assertEqual( + channel.json_body["errcode"], Codes.AS_FEDPROXY_DESTINATION_DENIED + ) + self.agent_request.assert_not_called() + + def test_self_destination_is_rejected(self) -> None: + channel = self._fed_proxy( + { + "destination": self.hs.hostname, + "method": "GET", + "path": f"/_matrix/federation/{APPSERVICE_PREFIX}/some/path", + } + ) + + self.assertEqual(channel.code, 403) + self.assertEqual( + channel.json_body["errcode"], Codes.AS_FEDPROXY_DESTINATION_DENIED + ) + self.agent_request.assert_not_called() + + def test_path_traversal_segment_is_rejected(self) -> None: + channel = self._fed_proxy( + { + "destination": "remote.example.com", + "method": "GET", + "path": f"/_matrix/federation/{APPSERVICE_PREFIX}/../../v1/send/txn1", + } + ) + + self.assertEqual(channel.code, 403) + self.assertEqual( + channel.json_body["errcode"], Codes.AS_FEDPROXY_PATH_NOT_ALLOWED + ) + self.agent_request.assert_not_called() + + def test_get_with_body_is_rejected(self) -> None: + channel = self._fed_proxy( + { + "destination": "remote.example.com", + "method": "GET", + "path": f"/_matrix/federation/{APPSERVICE_PREFIX}/some/path", + "body": {"key": "value"}, + } + ) + + self.assertEqual(channel.code, 400) + self.assertEqual(channel.json_body["errcode"], Codes.INVALID_PARAM) + self.agent_request.assert_not_called() + + def test_delete_with_body_is_rejected(self) -> None: + channel = self._fed_proxy( + { + "destination": "remote.example.com", + "method": "DELETE", + "path": f"/_matrix/federation/{APPSERVICE_PREFIX}/some/path", + "body": {"key": "value"}, + } + ) + + self.assertEqual(channel.code, 400) + self.assertEqual(channel.json_body["errcode"], Codes.INVALID_PARAM) + self.agent_request.assert_not_called() + + def test_non_string_query_value_is_rejected(self) -> None: + channel = self._fed_proxy( + { + "destination": "remote.example.com", + "method": "GET", + "path": f"/_matrix/federation/{APPSERVICE_PREFIX}/some/path", + "query": {"active": True}, + } + ) + + self.assertEqual(channel.code, 400) + self.assertEqual(channel.json_body["errcode"], Codes.INVALID_PARAM) + self.agent_request.assert_not_called() + + def test_path_outside_prefix_is_rejected(self) -> None: + channel = self._fed_proxy( + { + "destination": "remote.example.com", + "method": "GET", + "path": "/_matrix/federation/v1/send/txnid", + } + ) + + self.assertEqual(channel.code, 403) + self.assertEqual( + channel.json_body["errcode"], Codes.AS_FEDPROXY_PATH_NOT_ALLOWED + ) + self.agent_request.assert_not_called() + + def test_appservice_without_proxy_prefix_is_rejected(self) -> None: + other_token = "other_as_token" + other_appservice = ApplicationService( + other_token, + id="other_as", + sender=UserID.from_string("@other_bot:test"), + namespaces={}, + url=APPSERVICE_URL, + ) + self.hs.get_datastores().main.services_cache.append(other_appservice) + + channel = self._fed_proxy( + { + "destination": "remote.example.com", + "method": "GET", + "path": f"/_matrix/federation/{APPSERVICE_PREFIX}/some/path", + }, + access_token=other_token, + ) + + self.assertEqual(channel.code, 400) + self.assertEqual( + channel.json_body["errcode"], Codes.AS_FEDPROXY_NO_PROXY_PREFIX + ) + self.agent_request.assert_not_called() + + def test_unauthenticated_request_is_rejected(self) -> None: + channel = self._fed_proxy( + { + "destination": "remote.example.com", + "method": "GET", + "path": f"/_matrix/federation/{APPSERVICE_PREFIX}/some/path", + }, + access_token=None, + ) + + self.assertEqual(channel.code, 401) + self.agent_request.assert_not_called() + + def test_non_appservice_token_is_rejected(self) -> None: + self.register_user("normal_user", "password") + user_token = self.login("normal_user", "password") + + channel = self._fed_proxy( + { + "destination": "remote.example.com", + "method": "GET", + "path": f"/_matrix/federation/{APPSERVICE_PREFIX}/some/path", + }, + access_token=user_token, + ) + + self.assertEqual(channel.code, 403) + self.agent_request.assert_not_called() + + @unittest.override_config({"experimental_features": {"msc4512_enabled": False}}) + def test_endpoint_not_registered_when_msc4512_disabled(self) -> None: + channel = self._fed_proxy( + { + "destination": "remote.example.com", + "method": "GET", + "path": f"/_matrix/federation/{APPSERVICE_PREFIX}/some/path", + } + ) + + self.assertEqual(channel.code, 404) + self.agent_request.assert_not_called() diff --git a/tests/rest/client/test_appservice_proxy.py b/tests/rest/client/test_appservice_proxy.py new file mode 100644 index 00000000000..420f012f8e9 --- /dev/null +++ b/tests/rest/client/test_appservice_proxy.py @@ -0,0 +1,348 @@ +# +# 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.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 = "unstable/io.element.msc4195/rtc/livekit" + + +class ApplicationServiceClientProxyTestCase(unittest.HomeserverTestCase): + servlets = [ + admin.register_servlets, + login.register_servlets, + appservice_proxy.register_servlets, + ] + + def default_config(self) -> JsonDict: + config = super().default_config() + _, path = tempfile.mkstemp(prefix="as_proxy_config") + with open(path, "w") as f: + yaml.dump( + { + "id": "proxy_as", + "url": APPSERVICE_URL, + "as_token": "as_token", + "hs_token": "hs_token", + "sender_localpart": "proxy_bot", + "namespaces": {}, + "io.element.msc4512.proxy": APPSERVICE_PREFIX, + }, + 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: + 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/{APPSERVICE_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/{APPSERVICE_PREFIX}/some/path?foo=bar".encode(), + ) + + headers: Headers = kwargs["headers"] + self.assertIsNone(headers.getRawHeaders(b"Authorization")) + self.assertEqual( + headers.getRawHeaders(b"X-Matrix-User-Identifier"), + [self.user_id.encode("ascii")], + ) + + 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/{APPSERVICE_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/{APPSERVICE_PREFIX}".encode(), + ) + + headers: Headers = kwargs["headers"] + self.assertIsNone(headers.getRawHeaders(b"Authorization")) + 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/{APPSERVICE_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/{APPSERVICE_PREFIX}/some/path".encode(), + ) + + headers: Headers = kwargs["headers"] + self.assertIsNone(headers.getRawHeaders(b"Authorization")) + self.assertEqual( + headers.getRawHeaders(b"X-Matrix-User-Identifier"), + [self.user_id.encode("ascii")], + ) + + body_producer = kwargs["bodyProducer"] + expected_body = json_encoder.encode({"key": "value"}).encode("utf8") + self.assertEqual(body_producer.length, len(expected_body)) + + def test_connection_header_not_forwarded(self) -> None: + self.agent.request = Mock( + return_value=defer.succeed(FakeResponse.json(code=200, payload={})) + ) + + self.make_request( + "GET", + f"/_matrix/client/{APPSERVICE_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.assertEqual(headers.getRawHeaders(b"X-Forward"), [b"forward"]) + + def test_connection_header_with_named_request_headers_not_forwarded(self) -> None: + self.agent.request = Mock( + return_value=defer.succeed(FakeResponse.json(code=200, payload={})) + ) + + self.make_request( + "GET", + f"/_matrix/client/{APPSERVICE_PREFIX}/some/path", + shorthand=False, + access_token=self.access_token, + custom_headers=[ + ("Connection", "close, X-Omit"), + ("X-Omit", "omit"), + ("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-Omit")) + self.assertEqual(headers.getRawHeaders(b"X-Forward"), [b"forward"]) + + 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/{APPSERVICE_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/{APPSERVICE_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/{APPSERVICE_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/{APPSERVICE_PREFIX}/does/not/exist", + shorthand=False, + access_token=self.access_token, + ) + + self.assertEqual(channel.code, 404) + 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/{APPSERVICE_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/{APPSERVICE_PREFIX}/some/path", + shorthand=False, + access_token=self.access_token, + ) + + self.assertEqual(channel.code, 429) + 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/{APPSERVICE_PREFIX}-2", + shorthand=False, + ) + + self.assertEqual(channel.code, 404) + + @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/{APPSERVICE_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 4b9d069d6a3..c85f45de5c3 100644 --- a/tests/storage/test_appservice.py +++ b/tests/storage/test_appservice.py @@ -21,7 +21,7 @@ import json import os import tempfile -from typing import cast +from typing import Any, cast from unittest.mock import AsyncMock, Mock import yaml @@ -479,8 +479,8 @@ def __init__( class ApplicationServiceStoreConfigTestCase(unittest.HomeserverTestCase): - def _write_config(self, suffix: str, **kwargs: str) -> str: - vals = { + def _write_config(self, suffix: str, **kwargs: str | None) -> str: + vals: dict[str, Any] = { "id": "id" + suffix, "url": "url" + suffix, "as_token": "as_token" + suffix, @@ -566,3 +566,144 @@ def test_duplicate_as_tokens(self) -> None: self.assertIn(f1, str(e)) self.assertIn(f2, str(e)) self.assertIn("as_token", str(e)) + + def test_proxy_prefix_works(self) -> None: + f1 = self._write_config( + suffix="1", + **{ + "io.element.msc4512.proxy": "unstable/io.element.msc4195/rtc/livekit", + "url": "http://url1", + }, + ) + + 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, "unstable/io.element.msc4195/rtc/livekit" + ) + + def test_proxy_prefix_requires_url(self) -> None: + f1 = self._write_config( + suffix="1", + **{ + "io.element.msc4512.proxy": "unstable/io.element.msc4195/rtc/livekit", + "url": None, + }, + ) + + 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_url(self) -> None: + f1 = self._write_config( + suffix="1", + **{ + "io.element.msc4512.proxy": "unstable/io.element.msc4195/rtc/livekit", + "url": "", + }, + ) + + 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_does_not_allow_reserved_values(self) -> None: + f1 = self._write_config( + suffix="1", **{"io.element.msc4512.proxy": "not/allowed", "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_duplicate_proxy_prefix(self) -> None: + f1 = self._write_config( + suffix="1", + **{ + "io.element.msc4512.proxy": "unstable/io.element.msc4195/rtc/livekit", + "url": "http://url1", + }, + ) + f2 = self._write_config( + suffix="2", + **{ + "io.element.msc4512.proxy": "unstable/io.element.msc4195/rtc/livekit", + "url": "http://url2", + }, + ) + + 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", str(e))