Fix mypy: keep typed references to test mocks

This commit is contained in:
Matthew Hodgson
2026-07-23 18:55:49 +01:00
parent 9bc370fbe8
commit 9f06969efc
2 changed files with 32 additions and 29 deletions
+5 -4
View File
@@ -106,8 +106,8 @@ class FederationThirdPartyLookupTests(unittest.FederatingHomeserverTestCase):
def prepare(self, reactor: MemoryReactor, clock: Clock, hs: HomeServer) -> None:
super().prepare(reactor, clock, hs)
self.appservice_handler = hs.get_application_service_handler()
self.appservice_handler.query_3pe = AsyncMock( # type: ignore[method-assign]
appservice_handler = hs.get_application_service_handler()
self.query_3pe_mock = AsyncMock(
return_value=[
{
"userid": "@_xmpp_someone:test",
@@ -116,7 +116,8 @@ class FederationThirdPartyLookupTests(unittest.FederatingHomeserverTestCase):
}
]
)
self.appservice_handler.get_3pe_protocols = AsyncMock( # type: ignore[method-assign]
appservice_handler.query_3pe = self.query_3pe_mock # type: ignore[method-assign]
appservice_handler.get_3pe_protocols = AsyncMock( # type: ignore[method-assign]
return_value={"xmpp": {"instances": []}}
)
@@ -128,7 +129,7 @@ class FederationThirdPartyLookupTests(unittest.FederatingHomeserverTestCase):
self.assertEqual(
channel.json_body["results"][0]["userid"], "@_xmpp_someone:test"
)
call = self.appservice_handler.query_3pe.call_args
call = self.query_3pe_mock.call_args
self.assertEqual(call.args[1], "xmpp")
self.assertEqual(call.args[2], {b"username": [b"someone"]})
+27 -25
View File
@@ -58,21 +58,23 @@ class ThirdPartyFederatedLookupTests(unittest.HomeserverTestCase):
return config
def prepare(self, reactor: MemoryReactor, clock: Clock, hs: HomeServer) -> None:
self.appservice_handler = hs.get_application_service_handler()
self.appservice_handler.query_3pe = AsyncMock( # type: ignore[method-assign]
return_value=[LOCAL_USER_RESULT]
)
self.appservice_handler.get_3pe_protocols = AsyncMock( # type: ignore[method-assign]
appservice_handler = hs.get_application_service_handler()
self.query_3pe_mock = AsyncMock(return_value=[LOCAL_USER_RESULT])
appservice_handler.query_3pe = self.query_3pe_mock # type: ignore[method-assign]
self.get_3pe_protocols_mock = AsyncMock(
return_value={"xmpp": {"instances": [{"desc": "local xmpp"}]}}
)
appservice_handler.get_3pe_protocols = self.get_3pe_protocols_mock # type: ignore[method-assign]
self.transport = hs.get_federation_client().transport_layer
self.transport.get_thirdparty_entities = AsyncMock( # type: ignore[method-assign]
transport = hs.get_federation_client().transport_layer
self.get_thirdparty_entities_mock = AsyncMock(
return_value={"results": [REMOTE_USER_RESULT]}
)
self.transport.get_thirdparty_protocols = AsyncMock( # type: ignore[method-assign]
transport.get_thirdparty_entities = self.get_thirdparty_entities_mock # type: ignore[method-assign]
self.get_thirdparty_protocols_mock = AsyncMock(
return_value={"xmpp": {"instances": [{"desc": "remote xmpp"}]}}
)
transport.get_thirdparty_protocols = self.get_thirdparty_protocols_mock # type: ignore[method-assign]
self.user = self.register_user("user", "pass")
self.token = self.login(self.user, "pass")
@@ -85,7 +87,7 @@ class ThirdPartyFederatedLookupTests(unittest.HomeserverTestCase):
)
self.assertEqual(channel.code, 200, channel.result)
self.assertEqual(channel.json_list, [LOCAL_USER_RESULT])
self.transport.get_thirdparty_entities.assert_not_called()
self.get_thirdparty_entities_mock.assert_not_called()
def test_remote_user_lookup(self) -> None:
channel = self.make_request(
@@ -96,9 +98,9 @@ class ThirdPartyFederatedLookupTests(unittest.HomeserverTestCase):
)
self.assertEqual(channel.code, 200, channel.result)
self.assertEqual(channel.json_list, [REMOTE_USER_RESULT])
self.appservice_handler.query_3pe.assert_not_called()
self.query_3pe_mock.assert_not_called()
call = self.transport.get_thirdparty_entities.call_args
call = self.get_thirdparty_entities_mock.call_args
self.assertEqual(call.args[0], "remote.example.com")
self.assertEqual(call.args[1], "user")
self.assertEqual(call.args[2], "xmpp")
@@ -106,7 +108,7 @@ class ThirdPartyFederatedLookupTests(unittest.HomeserverTestCase):
self.assertEqual(call.args[3], {"username": ["someone"]})
def test_remote_location_lookup(self) -> None:
self.transport.get_thirdparty_entities.return_value = {
self.get_thirdparty_entities_mock.return_value = {
"results": [REMOTE_LOCATION_RESULT]
}
channel = self.make_request(
@@ -118,7 +120,7 @@ class ThirdPartyFederatedLookupTests(unittest.HomeserverTestCase):
self.assertEqual(channel.code, 200, channel.result)
self.assertEqual(channel.json_list, [REMOTE_LOCATION_RESULT])
self.assertEqual(
self.transport.get_thirdparty_entities.call_args.args[1], "location"
self.get_thirdparty_entities_mock.call_args.args[1], "location"
)
def test_remote_protocols(self) -> None:
@@ -131,7 +133,7 @@ class ThirdPartyFederatedLookupTests(unittest.HomeserverTestCase):
self.assertEqual(
channel.json_body, {"xmpp": {"instances": [{"desc": "remote xmpp"}]}}
)
self.appservice_handler.get_3pe_protocols.assert_not_called()
self.get_3pe_protocols_mock.assert_not_called()
def test_own_server_name_is_local(self) -> None:
channel = self.make_request(
@@ -141,13 +143,13 @@ class ThirdPartyFederatedLookupTests(unittest.HomeserverTestCase):
)
self.assertEqual(channel.code, 200, channel.result)
self.assertEqual(channel.json_list, [LOCAL_USER_RESULT])
self.transport.get_thirdparty_entities.assert_not_called()
self.get_thirdparty_entities_mock.assert_not_called()
# The stripped server param must not reach the local bridges either.
fields = self.appservice_handler.query_3pe.call_args.args[2]
fields = self.query_3pe_mock.call_args.args[2]
self.assertNotIn(b"server", fields)
def test_remote_failure_returns_empty(self) -> None:
self.transport.get_thirdparty_entities.side_effect = RuntimeError("boom")
self.get_thirdparty_entities_mock.side_effect = RuntimeError("boom")
channel = self.make_request(
"GET",
"/_matrix/client/v3/thirdparty/user/xmpp"
@@ -168,14 +170,14 @@ class ThirdPartyFederatedLookupDisabledTests(unittest.HomeserverTestCase):
]
def prepare(self, reactor: MemoryReactor, clock: Clock, hs: HomeServer) -> None:
self.appservice_handler = hs.get_application_service_handler()
self.appservice_handler.query_3pe = AsyncMock( # type: ignore[method-assign]
return_value=[LOCAL_USER_RESULT]
)
self.transport = hs.get_federation_client().transport_layer
self.transport.get_thirdparty_entities = AsyncMock( # type: ignore[method-assign]
appservice_handler = hs.get_application_service_handler()
self.query_3pe_mock = AsyncMock(return_value=[LOCAL_USER_RESULT])
appservice_handler.query_3pe = self.query_3pe_mock # type: ignore[method-assign]
transport = hs.get_federation_client().transport_layer
self.get_thirdparty_entities_mock = AsyncMock(
return_value={"results": [REMOTE_USER_RESULT]}
)
transport.get_thirdparty_entities = self.get_thirdparty_entities_mock # type: ignore[method-assign]
self.user = self.register_user("user", "pass")
self.token = self.login(self.user, "pass")
@@ -189,7 +191,7 @@ class ThirdPartyFederatedLookupDisabledTests(unittest.HomeserverTestCase):
)
self.assertEqual(channel.code, 200, channel.result)
self.assertEqual(channel.json_list, [LOCAL_USER_RESULT])
self.transport.get_thirdparty_entities.assert_not_called()
self.get_thirdparty_entities_mock.assert_not_called()
# The server param is still stripped from the bridge query fields.
fields = self.appservice_handler.query_3pe.call_args.args[2]
fields = self.query_3pe_mock.call_args.args[2]
self.assertNotIn(b"server", fields)