| 1 | from __future__ import annotations |
| 2 | |
| 3 | from datetime import datetime, timezone |
| 4 | |
| 5 | import pytest |
| 6 | |
| 7 | from app.contracts import McpErrorCode, McpGatewayError, McpRiskClass |
| 8 | from app.external import ( |
| 9 | ExternalMcpGateway, |
| 10 | ExternalMcpProviderMetadata, |
| 11 | McpServerRegistry, |
| 12 | ProviderFallbackAuthorization, |
| 13 | ) |
| 14 | from conftest import INSTRUMENT_ID |
| 15 | |
| 16 | |
| 17 | class MockExternalProvider: |
| 18 | metadata = ExternalMcpProviderMetadata( |
| 19 | providerId="TEST_MCP", |
| 20 | regions=("INDIA",), |
| 21 | supportedRequirements=("CURRENT_NEWS",), |
| 22 | supportedTools=("SEARCH_NEWS",), |
| 23 | riskClass=McpRiskClass.SAFE_READ, |
| 24 | authType="MOCK", |
| 25 | ) |
| 26 | |
| 27 | def __init__(self, *, supported=True, result=None) -> None: |
| 28 | self.supported = supported |
| 29 | self.result = result or {"globalInstrumentId": str(INSTRUMENT_ID), "evidence": []} |
| 30 | self.calls = 0 |
| 31 | |
| 32 | async def supports(self, **_kwargs) -> bool: |
| 33 | return self.supported |
| 34 | |
| 35 | async def invoke(self, **_kwargs): |
| 36 | self.calls += 1 |
| 37 | return self.result |
| 38 | |
| 39 | async def health(self): |
| 40 | return {"status": "UP"} |
| 41 | |
| 42 | |
| 43 | def authorization(**overrides) -> ProviderFallbackAuthorization: |
| 44 | values = { |
| 45 | "authorized": True, |
| 46 | "issuedBy": "ProviderFallbackPolicy", |
| 47 | "globalInstrumentId": INSTRUMENT_ID, |
| 48 | "requirementId": "CURRENT_NEWS", |
| 49 | "permittedProviderIds": ("TEST_MCP",), |
| 50 | "issuedAt": datetime.now(timezone.utc), |
| 51 | } |
| 52 | values.update(overrides) |
| 53 | return ProviderFallbackAuthorization(**values) |
| 54 | |
| 55 | |
| 56 | def test_mock_external_provider_can_register_and_duplicates_fail() -> None: |
| 57 | registry = McpServerRegistry() |
| 58 | provider = MockExternalProvider() |
| 59 | registry.register(provider) |
| 60 | assert registry.get("test_mcp") is provider |
| 61 | assert provider.metadata.kind.value == "EXTERNAL" |
| 62 | with pytest.raises(ValueError, match="already registered"): |
| 63 | registry.register(provider) |
| 64 | |
| 65 | |
| 66 | @pytest.mark.asyncio |
| 67 | async def test_unsupported_requirement_is_rejected() -> None: |
| 68 | registry = McpServerRegistry() |
| 69 | provider = MockExternalProvider(supported=False) |
| 70 | registry.register(provider) |
| 71 | gateway = ExternalMcpGateway(registry, enabled=True) |
| 72 | with pytest.raises(McpGatewayError) as error: |
| 73 | await gateway.invoke( |
| 74 | provider_id="TEST_MCP", |
| 75 | region="INDIA", |
| 76 | requirement_id="CURRENT_NEWS", |
| 77 | tool="SEARCH_NEWS", |
| 78 | global_instrument_id=INSTRUMENT_ID, |
| 79 | arguments={}, |
| 80 | request_id="external-5a", |
| 81 | authorization=authorization(), |
| 82 | ) |
| 83 | assert error.value.code == McpErrorCode.EXTERNAL_PROVIDER_UNAVAILABLE |
| 84 | assert provider.calls == 0 |
| 85 | |
| 86 | |
| 87 | @pytest.mark.asyncio |
| 88 | async def test_fallback_requires_explicit_policy_authorization() -> None: |
| 89 | registry = McpServerRegistry() |
| 90 | provider = MockExternalProvider() |
| 91 | registry.register(provider) |
| 92 | gateway = ExternalMcpGateway(registry, enabled=True) |
| 93 | with pytest.raises(McpGatewayError) as error: |
| 94 | await gateway.invoke( |
| 95 | provider_id="TEST_MCP", |
| 96 | region="INDIA", |
| 97 | requirement_id="CURRENT_NEWS", |
| 98 | tool="SEARCH_NEWS", |
| 99 | global_instrument_id=INSTRUMENT_ID, |
| 100 | arguments={}, |
| 101 | request_id="external-5a", |
| 102 | authorization=None, |
| 103 | ) |
| 104 | assert error.value.code == McpErrorCode.FORBIDDEN |
| 105 | assert provider.calls == 0 |
| 106 | |
| 107 | |
| 108 | @pytest.mark.asyncio |
| 109 | async def test_external_provider_cannot_create_or_change_canonical_identity() -> None: |
| 110 | for result in ( |
| 111 | {"createIdentity": {"symbol": "NEW"}}, |
| 112 | {"globalInstrumentId": "99999999-9999-4999-8999-999999999999"}, |
| 113 | {"nested": {"providerMappings": [{"symbol": "NEW"}]}}, |
| 114 | ): |
| 115 | registry = McpServerRegistry() |
| 116 | provider = MockExternalProvider(result=result) |
| 117 | registry.register(provider) |
| 118 | gateway = ExternalMcpGateway(registry, enabled=True) |
| 119 | with pytest.raises(McpGatewayError) as error: |
| 120 | await gateway.invoke( |
| 121 | provider_id="TEST_MCP", |
| 122 | region="INDIA", |
| 123 | requirement_id="CURRENT_NEWS", |
| 124 | tool="SEARCH_NEWS", |
| 125 | global_instrument_id=INSTRUMENT_ID, |
| 126 | arguments={}, |
| 127 | request_id="external-5a", |
| 128 | authorization=authorization(), |
| 129 | ) |
| 130 | assert error.value.code == McpErrorCode.FORBIDDEN |
| 131 | |
| 132 | |
| 133 | @pytest.mark.asyncio |
| 134 | async def test_policy_authorized_external_result_keeps_canonical_identity_and_is_sanitized() -> None: |
| 135 | registry = McpServerRegistry() |
| 136 | provider = MockExternalProvider( |
| 137 | result={ |
| 138 | "globalInstrumentId": str(INSTRUMENT_ID), |
| 139 | "evidence": [{"title": "safe", "apiKey": "must-not-leak"}], |
| 140 | } |
| 141 | ) |
| 142 | registry.register(provider) |
| 143 | gateway = ExternalMcpGateway(registry, enabled=True) |
| 144 | result = await gateway.invoke( |
| 145 | provider_id="TEST_MCP", |
| 146 | region="INDIA", |
| 147 | requirement_id="CURRENT_NEWS", |
| 148 | tool="SEARCH_NEWS", |
| 149 | global_instrument_id=INSTRUMENT_ID, |
| 150 | arguments={}, |
| 151 | request_id="external-allowed-5a", |
| 152 | authorization=authorization(), |
| 153 | ) |
| 154 | assert result == { |
| 155 | "globalInstrumentId": str(INSTRUMENT_ID), |
| 156 | "evidence": [{"title": "safe"}], |
| 157 | } |
| 158 | assert provider.calls == 1 |
| 159 | |
| 160 | |
| 161 | @pytest.mark.asyncio |
| 162 | async def test_external_provider_exception_is_normalized() -> None: |
| 163 | class FailedProvider(MockExternalProvider): |
| 164 | async def invoke(self, **_kwargs): |
| 165 | raise RuntimeError("provider credential and stack must not escape") |
| 166 | |
| 167 | registry = McpServerRegistry() |
| 168 | registry.register(FailedProvider()) |
| 169 | gateway = ExternalMcpGateway(registry, enabled=True) |
| 170 | with pytest.raises(McpGatewayError) as error: |
| 171 | await gateway.invoke( |
| 172 | provider_id="TEST_MCP", |
| 173 | region="INDIA", |
| 174 | requirement_id="CURRENT_NEWS", |
| 175 | tool="SEARCH_NEWS", |
| 176 | global_instrument_id=INSTRUMENT_ID, |
| 177 | arguments={}, |
| 178 | request_id="external-failure-5a", |
| 179 | authorization=authorization(), |
| 180 | ) |
| 181 | assert error.value.code == McpErrorCode.EXTERNAL_PROVIDER_UNAVAILABLE |
| 182 | assert "credential" not in str(error.value) |
| 183 | |
| 184 | |
| 185 | @pytest.mark.asyncio |
| 186 | async def test_non_safe_external_provider_is_denied_before_invocation() -> None: |
| 187 | provider = MockExternalProvider() |
| 188 | provider.metadata = provider.metadata.model_copy( |
| 189 | update={"risk_class": McpRiskClass.FINANCIAL_ACTION} |
| 190 | ) |
| 191 | registry = McpServerRegistry() |
| 192 | registry.register(provider) |
| 193 | gateway = ExternalMcpGateway(registry, enabled=True) |
| 194 | with pytest.raises(McpGatewayError) as error: |
| 195 | await gateway.invoke( |
| 196 | provider_id="TEST_MCP", |
| 197 | region="INDIA", |
| 198 | requirement_id="CURRENT_NEWS", |
| 199 | tool="SEARCH_NEWS", |
| 200 | global_instrument_id=INSTRUMENT_ID, |
| 201 | arguments={}, |
| 202 | request_id="external-denied-5a", |
| 203 | authorization=authorization(), |
| 204 | ) |
| 205 | assert error.value.code == McpErrorCode.MCP_TOOL_DENIED |
| 206 | assert provider.calls == 0 |