Files
synapse/tests/rest/client/test_appservice_proxy.py
T
Johannes Marbach 57d6da409c Add experimental support for letting application services proxy namespaces in the C-S and S-S API as per MSC4512 (#19972)
This implements the proxying part of
[MSC4512](https://github.com/matrix-org/matrix-spec-proposals/pull/4512)
and is a stopgap towards
https://github.com/element-hq/voip-internal/issues/641. It introduces a
new configuration property `io.element.msc4512.proxy` that allows
application services to claim namespaces in the C-S and S-S API. For
requests underneath a claimed namespace, Synapse first authorizes the
request and then reverse-proxies it to the application services. For
now, the only allowed namespace that can be claimed is
`unstable/io.element.msc4195/rtc/livekit`.

This pull request can be reviewed by commits.

### Pull Request Checklist

<!-- Please read
https://element-hq.github.io/synapse/latest/development/contributing_guide.html
before submitting your pull request -->

* [x] Pull request is based on the develop branch
* [x] Pull request includes a [changelog
file](https://element-hq.github.io/synapse/latest/development/contributing_guide.html#changelog).
The entry should:
- Be a short description of your change which makes sense to users.
"Fixed a bug that prevented receiving messages from other servers."
instead of "Moved X method from `EventStore` to `EventWorkerStore`.".
  - Use markdown where necessary, mostly for `code blocks`.
  - End with either a period (.) or an exclamation mark (!).
  - Start with a capital letter.
- Feel free to credit yourself, by adding a sentence "Contributed by
@github_username." or "Contributed by [Your Name]." to the end of the
entry.
* [x] [Code
style](https://element-hq.github.io/synapse/latest/code_style.html) is
correct (run the
[linters](https://element-hq.github.io/synapse/latest/development/contributing_guide.html#run-the-linters))

---------

Signed-off-by: Johannes Marbach <n0-0ne+github@mailbox.org>
2026-08-28 13:45:54 +01:00

443 lines
15 KiB
Python

#
# 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:
# <https://www.gnu.org/licenses/agpl-3.0.html>.
#
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()