From 5e10ea41365b56dd3ada20d2189d9f52af7a7320 Mon Sep 17 00:00:00 2001 From: Talal Date: Wed, 29 Oct 2025 19:11:32 -0700 Subject: [PATCH] Improve(mcp): respect X-Forwarded- headers in OAuth endpoints (#16036) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix(mcp): respect X-Forwarded-Proto header in OAuth endpoints When LiteLLM proxy is deployed behind a reverse proxy (like nginx or a load balancer) that terminates SSL/TLS, the proxy receives HTTP requests internally but should expose HTTPS URLs externally. This change detects the X-Forwarded-Proto header and uses it to construct correct redirect URIs and endpoint URLs. Changes: - Added X-Forwarded-Proto detection to authorize, token, oauth_protected_resource_mcp, oauth_authorization_server_mcp, and register_client endpoints - Added comprehensive tests for X-Forwarded-Proto header support across all affected endpoints - Fixed existing tests to properly mock request.headers 🤖 Generated with [Claude Code](https://claude.com/claude-code) Co-Authored-By: Claude * fix formatting * feat(mcp): support X-Forwarded-Host for proxy base URL reconstruction Extended X-Forwarded-Proto support to also handle X-Forwarded-Host and X-Forwarded-Port headers. This allows LiteLLM to correctly construct redirect URIs and endpoint URLs when deployed behind a reverse proxy that changes the host/port. Example scenario: - Internal URL: http://localhost:8888/github/mcp - External URL: https://proxy.abc.com/github/mcp - Proxy sets: X-Forwarded-Proto: https, X-Forwarded-Host: proxy.abc.com Changes: - Added get_request_base_url() helper function to centralize X-Forwarded-* header handling - Replaced all inline X-Forwarded-Proto checks with calls to the helper function - Helper handles X-Forwarded-Proto, X-Forwarded-Host, and X-Forwarded-Port - Added tests for X-Forwarded-Host scenarios in authorize and token endpoints Fixes issue where protected resource URL mismatch occurred: Error: Protected resource http://proxy.abc.com:8888/github/mcp does not match expected https://proxy.abc.com/github/mcp 🤖 Generated with [Claude Code](https://claude.com/claude-code) Co-Authored-By: Claude * chore: replace Yelp-specific hostnames with generic examples Changed all references from chatproxy.yelpcorp.com to proxy.example.com in: - test_proxy_forwarding.py (default host parameter) - TEST_PROXY_FORWARDING.md (documentation examples) - discoverable_endpoints.py (docstring example) - test_discoverable_endpoints.py (test mock data) This makes the code more generic and suitable for open source. All 13 tests still passing. * remove accidentally added files * fix formatting * add new test for get_base_url --------- Co-authored-by: Claude --- .../mcp_server/discoverable_endpoints.py | 65 +- .../mcp_server/test_discoverable_endpoints.py | 576 +++++++++++++++++- 2 files changed, 629 insertions(+), 12 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 59a443f7b79..583c83cca51 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -20,6 +20,55 @@ router = APIRouter( ) +def get_request_base_url(request: Request) -> str: + """ + Get the base URL for the request, considering X-Forwarded-* headers. + + When behind a proxy (like nginx), the proxy may set: + - X-Forwarded-Proto: The original protocol (http/https) + - X-Forwarded-Host: The original host (may include port) + - X-Forwarded-Port: The original port (if not in Host header) + + Args: + request: FastAPI Request object + + Returns: + The reconstructed base URL (e.g., "https://proxy.example.com") + """ + base_url = str(request.base_url).rstrip("/") + parsed = urlparse(base_url) + + # Get forwarded headers + x_forwarded_proto = request.headers.get("X-Forwarded-Proto") + x_forwarded_host = request.headers.get("X-Forwarded-Host") + x_forwarded_port = request.headers.get("X-Forwarded-Port") + + # Start with the original scheme + scheme = x_forwarded_proto if x_forwarded_proto else parsed.scheme + + # Handle host and port + if x_forwarded_host: + # X-Forwarded-Host may already include port (e.g., "example.com:8080") + if ":" in x_forwarded_host and not x_forwarded_host.startswith("["): + # Host includes port + netloc = x_forwarded_host + elif x_forwarded_port: + # Port is separate + netloc = f"{x_forwarded_host}:{x_forwarded_port}" + else: + # Just host, no explicit port + netloc = x_forwarded_host + else: + # No X-Forwarded-Host, use original netloc + netloc = parsed.netloc + if x_forwarded_port and ":" not in netloc: + # Add forwarded port if not already in netloc + netloc = f"{netloc}:{x_forwarded_port}" + + # Reconstruct the URL + return urlunparse((scheme, netloc, parsed.path, "", "", "")) + + def encode_state_with_base_url( base_url: str, original_state: str, @@ -107,7 +156,9 @@ async def authorize( # Parse it to remove any existing query parsed = urlparse(redirect_uri) base_url = urlunparse(parsed._replace(query="")) - request_base_url = str(request.base_url).rstrip("/") + + # Get the correct base URL considering X-Forwarded-* headers + request_base_url = get_request_base_url(request) # Encode the base_url, original state, PKCE params, and client redirect_uri in encrypted state encoded_state = encode_state_with_base_url( @@ -177,7 +228,8 @@ async def token_endpoint( if mcp_server.token_url is None: raise HTTPException(status_code=400, detail="MCP server token url is not set") - proxy_base_url = str(request.base_url).rstrip("/") + # Get the correct base URL considering X-Forwarded-* headers + proxy_base_url = get_request_base_url(request) # Build token request data token_data = { @@ -251,7 +303,8 @@ async def callback(code: str, state: str): async def oauth_protected_resource_mcp( request: Request, mcp_server_name: Optional[str] = None ): - request_base_url = str(request.base_url).rstrip("/") + # Get the correct base URL considering X-Forwarded-* headers + request_base_url = get_request_base_url(request) return { "authorization_servers": [ ( @@ -273,7 +326,8 @@ async def oauth_protected_resource_mcp( async def oauth_authorization_server_mcp( request: Request, mcp_server_name: Optional[str] = None ): - request_base_url = str(request.base_url).rstrip("/") + # Get the correct base URL considering X-Forwarded-* headers + request_base_url = get_request_base_url(request) authorization_endpoint = ( f"{request_base_url}/{mcp_server_name}/authorize" @@ -320,7 +374,8 @@ async def register_client(request: Request, mcp_server_name: Optional[str] = Non global_mcp_server_manager, ) - request_base_url = str(request.base_url).rstrip("/") + # Get the correct base URL considering X-Forwarded-* headers + request_base_url = get_request_base_url(request) request_data = await _read_request_body(request=request) data: dict = {**request_data} diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 8ed0dd0a12a..7afef2627ba 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -179,6 +179,7 @@ async def test_token_endpoint_forwards_code_verifier(): # Mock request mock_request = MagicMock(spec=Request) mock_request.base_url = "https://litellm-proxy.example.com/" + mock_request.headers = {} # Mock httpx client response mock_response = MagicMock() @@ -192,6 +193,7 @@ async def test_token_endpoint_forwards_code_verifier(): # Mock the async httpx client with AsyncMock for async methods from unittest.mock import AsyncMock + with patch( "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client" ) as mock_get_client: @@ -219,29 +221,35 @@ async def test_token_endpoint_forwards_code_verifier(): # Check the data parameter includes code_verifier assert call_args[1]["data"]["code_verifier"] == "test_code_verifier_from_client" assert call_args[1]["data"]["code"] == "4/test_authorization_code" - assert call_args[1]["data"]["client_id"] == "669428968603-test.apps.googleusercontent.com" + assert ( + call_args[1]["data"]["client_id"] + == "669428968603-test.apps.googleusercontent.com" + ) assert call_args[1]["data"]["client_secret"] == "GOCSPX-test_secret" assert call_args[1]["data"]["grant_type"] == "authorization_code" # Verify response response_data = response.body import json + token_data = json.loads(response_data) assert token_data["access_token"] == "ya29.test_access_token" assert token_data["token_type"] == "Bearer" - @pytest.mark.asyncio async def test_register_client_without_mcp_server_name_returns_dummy(): try: - from litellm.proxy._experimental.mcp_server.discoverable_endpoints import register_client + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + register_client, + ) from fastapi import Request except ImportError: pytest.skip("MCP discoverable endpoints not available") mock_request = MagicMock(spec=Request) mock_request.base_url = "https://proxy.litellm.example/" + mock_request.headers = {} with patch( "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", new=AsyncMock(return_value={}), @@ -258,8 +266,12 @@ async def test_register_client_without_mcp_server_name_returns_dummy(): @pytest.mark.asyncio async def test_register_client_returns_existing_server_credentials(): try: - from litellm.proxy._experimental.mcp_server.discoverable_endpoints import register_client - from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + register_client, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) from litellm.types.mcp import MCPAuth from litellm.types.mcp_server.mcp_server_manager import MCPServer from litellm.proxy._types import MCPTransport @@ -284,6 +296,7 @@ async def test_register_client_returns_existing_server_credentials(): mock_request = MagicMock(spec=Request) mock_request.base_url = "https://proxy.litellm.example/" + mock_request.headers = {} try: with patch( @@ -306,8 +319,12 @@ async def test_register_client_returns_existing_server_credentials(): @pytest.mark.asyncio async def test_register_client_remote_registration_success(): try: - from litellm.proxy._experimental.mcp_server.discoverable_endpoints import register_client - from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + register_client, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) from litellm.types.mcp import MCPAuth from litellm.types.mcp_server.mcp_server_manager import MCPServer from litellm.proxy._types import MCPTransport @@ -333,6 +350,7 @@ async def test_register_client_remote_registration_success(): mock_request = MagicMock(spec=Request) mock_request.base_url = "https://proxy.litellm.example/" + mock_request.headers = {} request_payload = { "client_name": "Litellm Proxy", @@ -387,3 +405,547 @@ async def test_register_client_remote_registration_success(): ) +@pytest.mark.asyncio +async def test_authorize_endpoint_respects_x_forwarded_proto(): + """Test that authorize endpoint uses X-Forwarded-Proto header to construct correct redirect_uri""" + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + authorize, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + from litellm.proxy._types import MCPTransport + from fastapi import Request + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + # Clear registry + global_mcp_server_manager.registry.clear() + + # Create mock OAuth2 server + oauth2_server = MCPServer( + server_id="test_oauth_server", + name="test_oauth", + server_name="test_oauth", + alias="test_oauth", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="test_client_id", + client_secret="test_client_secret", + authorization_url="https://provider.com/oauth/authorize", + token_url="https://provider.com/oauth/token", + scopes=["read", "write"], + ) + global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server + + # Mock request with http base_url but X-Forwarded-Proto: https + mock_request = MagicMock(spec=Request) + mock_request.base_url = "http://litellm.example.com/" # HTTP + mock_request.headers = {"X-Forwarded-Proto": "https"} # Behind HTTPS proxy + + # Mock the encryption functions + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper" + ) as mock_encrypt: + mock_encrypt.return_value = "mocked_encrypted_state" + + # Call authorize endpoint + response = await authorize( + request=mock_request, + client_id="test_client_id", + mcp_server_name="test_oauth", + redirect_uri="https://client.example.com/callback", + state="test_state", + ) + + # Verify redirect URL uses HTTPS in the redirect_uri parameter + location = response.headers["location"] + + # The redirect_uri parameter sent to the OAuth provider should use HTTPS + assert ( + "redirect_uri=https%3A%2F%2Flitellm.example.com%2Fcallback" in location + or "redirect_uri=https://litellm.example.com/callback" in location + ) + + +@pytest.mark.asyncio +async def test_token_endpoint_respects_x_forwarded_proto(): + """Test that token endpoint uses X-Forwarded-Proto header for redirect_uri""" + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + token_endpoint, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + from litellm.proxy._types import MCPTransport + from fastapi import Request + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + # Clear registry + global_mcp_server_manager.registry.clear() + + # Create mock OAuth2 server + oauth2_server = MCPServer( + server_id="google_mcp", + name="google_mcp", + server_name="google_mcp", + alias="google_mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="test_client_id", + client_secret="test_secret", + authorization_url="https://accounts.google.com/o/oauth2/v2/auth", + token_url="https://oauth2.googleapis.com/token", + scopes=["openid", "email"], + ) + global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server + + # Mock request with http base_url but X-Forwarded-Proto: https + mock_request = MagicMock(spec=Request) + mock_request.base_url = "http://litellm-proxy.example.com/" # HTTP + mock_request.headers = {"X-Forwarded-Proto": "https"} # Behind HTTPS proxy + + # Mock httpx client response + mock_response = MagicMock() + mock_response.json.return_value = { + "access_token": "test_token", + "token_type": "Bearer", + "expires_in": 3599, + } + mock_response.raise_for_status = MagicMock() + + # Mock the async httpx client + mock_async_client = MagicMock() + mock_async_client.post = AsyncMock(return_value=mock_response) + + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client" + ) as mock_get_client: + mock_get_client.return_value = mock_async_client + + # Call token endpoint + response = await token_endpoint( + request=mock_request, + grant_type="authorization_code", + code="test_code", + redirect_uri="http://localhost:60108/callback", + client_id="test_client_id", + mcp_server_name="google_mcp", + client_secret="test_secret", + ) + + # Verify that the redirect_uri sent to the provider uses HTTPS + call_args = mock_async_client.post.call_args + assert ( + call_args[1]["data"]["redirect_uri"] + == "https://litellm-proxy.example.com/callback" + ) + + +@pytest.mark.asyncio +async def test_oauth_protected_resource_respects_x_forwarded_proto(): + """Test that oauth_protected_resource_mcp uses X-Forwarded-Proto for URLs""" + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + oauth_protected_resource_mcp, + ) + from fastapi import Request + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + # Mock request with http base_url but X-Forwarded-Proto: https + mock_request = MagicMock(spec=Request) + mock_request.base_url = "http://litellm.example.com/" # HTTP + mock_request.headers = {"X-Forwarded-Proto": "https"} # Behind HTTPS proxy + + # Call the endpoint + response = await oauth_protected_resource_mcp( + request=mock_request, + mcp_server_name="test_server", + ) + + # Verify response uses HTTPS URLs + assert response["authorization_servers"][0].startswith( + "https://litellm.example.com/" + ) + + +@pytest.mark.asyncio +async def test_oauth_authorization_server_respects_x_forwarded_proto(): + """Test that oauth_authorization_server_mcp uses X-Forwarded-Proto for URLs""" + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + oauth_authorization_server_mcp, + ) + from fastapi import Request + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + # Mock request with http base_url but X-Forwarded-Proto: https + mock_request = MagicMock(spec=Request) + mock_request.base_url = "http://litellm.example.com/" # HTTP + mock_request.headers = {"X-Forwarded-Proto": "https"} # Behind HTTPS proxy + + # Call the endpoint + response = await oauth_authorization_server_mcp( + request=mock_request, + mcp_server_name="test_server", + ) + + # Verify response uses HTTPS URLs + assert response["authorization_endpoint"].startswith("https://litellm.example.com/") + assert response["token_endpoint"].startswith("https://litellm.example.com/") + assert response["registration_endpoint"].startswith("https://litellm.example.com/") + + +@pytest.mark.asyncio +async def test_register_client_respects_x_forwarded_proto(): + """Test that register_client uses X-Forwarded-Proto for redirect_uris""" + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + register_client, + ) + from fastapi import Request + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + # Mock request with http base_url but X-Forwarded-Proto: https + mock_request = MagicMock(spec=Request) + mock_request.base_url = "http://proxy.litellm.example/" # HTTP + mock_request.headers = {"X-Forwarded-Proto": "https"} # Behind HTTPS proxy + + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + new=AsyncMock(return_value={}), + ): + result = await register_client(request=mock_request) + + # Verify the redirect_uris use HTTPS + assert result == { + "client_id": "dummy_client", + "client_secret": "dummy", + "redirect_uris": ["https://proxy.litellm.example/callback"], + } + + +@pytest.mark.asyncio +async def test_authorize_endpoint_respects_x_forwarded_host(): + """Test that authorize endpoint uses X-Forwarded-Host and X-Forwarded-Proto to construct correct redirect_uri""" + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + authorize, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + from litellm.proxy._types import MCPTransport + from fastapi import Request + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + # Clear registry + global_mcp_server_manager.registry.clear() + + # Create mock OAuth2 server + oauth2_server = MCPServer( + server_id="test_oauth_server", + name="test_oauth", + server_name="test_oauth", + alias="test_oauth", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="test_client_id", + client_secret="test_client_secret", + authorization_url="https://provider.com/oauth/authorize", + token_url="https://provider.com/oauth/token", + scopes=["read", "write"], + ) + global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server + + # Mock request simulating nginx proxy: + # Internal: http://localhost:8888/github/mcp + # External: https://proxy.example.com/github/mcp + mock_request = MagicMock(spec=Request) + mock_request.base_url = "http://localhost:8888/github/mcp" + mock_request.headers = { + "X-Forwarded-Proto": "https", + "X-Forwarded-Host": "proxy.example.com", + } + + # Mock the encryption functions + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper" + ) as mock_encrypt: + mock_encrypt.return_value = "mocked_encrypted_state" + + # Call authorize endpoint + response = await authorize( + request=mock_request, + client_id="test_client_id", + mcp_server_name="test_oauth", + redirect_uri="https://client.example.com/callback", + state="test_state", + ) + + # Verify redirect URL uses the forwarded host and scheme + location = response.headers["location"] + + # The redirect_uri parameter should use the external URL + assert ( + "redirect_uri=https%3A%2F%2Fproxy.example.com%2Fgithub%2Fmcp%2Fcallback" + in location + or "redirect_uri=https://proxy.example.com/github/mcp/callback" in location + ) + + +@pytest.mark.asyncio +async def test_token_endpoint_respects_x_forwarded_host(): + """Test that token endpoint uses X-Forwarded-Host and X-Forwarded-Proto for redirect_uri""" + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + token_endpoint, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + from litellm.proxy._types import MCPTransport + from fastapi import Request + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + # Clear registry + global_mcp_server_manager.registry.clear() + + # Create mock OAuth2 server + oauth2_server = MCPServer( + server_id="google_mcp", + name="google_mcp", + server_name="google_mcp", + alias="google_mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="test_client_id", + client_secret="test_secret", + authorization_url="https://accounts.google.com/o/oauth2/v2/auth", + token_url="https://oauth2.googleapis.com/token", + scopes=["openid", "email"], + ) + global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server + + # Mock request simulating nginx proxy without port in host + mock_request = MagicMock(spec=Request) + mock_request.base_url = "http://localhost:8888/github/mcp" + mock_request.headers = { + "X-Forwarded-Proto": "https", + "X-Forwarded-Host": "proxy.example.com", + } + + # Mock httpx client response + mock_response = MagicMock() + mock_response.json.return_value = { + "access_token": "test_token", + "token_type": "Bearer", + "expires_in": 3599, + } + mock_response.raise_for_status = MagicMock() + + # Mock the async httpx client + mock_async_client = MagicMock() + mock_async_client.post = AsyncMock(return_value=mock_response) + + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client" + ) as mock_get_client: + mock_get_client.return_value = mock_async_client + + # Call token endpoint + response = await token_endpoint( + request=mock_request, + grant_type="authorization_code", + code="test_code", + redirect_uri="http://localhost:60108/callback", + client_id="test_client_id", + mcp_server_name="google_mcp", + client_secret="test_secret", + ) + + # Verify that the redirect_uri sent to the provider uses the external URL + call_args = mock_async_client.post.call_args + assert ( + call_args[1]["data"]["redirect_uri"] + == "https://proxy.example.com/github/mcp/callback" + ) + + +@pytest.mark.parametrize( + "base_url,x_forwarded_proto,x_forwarded_host,x_forwarded_port,expected_url", + [ + # Case 1: No forwarded headers - use original URL as-is (no trailing slash) + ( + "http://localhost:4000/", + None, + None, + None, + "http://localhost:4000", + ), + # Case 2: Only X-Forwarded-Proto - change scheme only + ( + "http://localhost:4000/", + "https", + None, + None, + "https://localhost:4000", + ), + # Case 3: X-Forwarded-Proto + X-Forwarded-Host - change scheme and host + ( + "http://localhost:4000/", + "https", + "proxy.example.com", + None, + "https://proxy.example.com", + ), + # Case 4: X-Forwarded-Host with port included in host header + ( + "http://localhost:4000/", + "https", + "proxy.example.com:8080", + None, + "https://proxy.example.com:8080", + ), + # Case 5: X-Forwarded-Host + X-Forwarded-Port as separate headers + ( + "http://localhost:4000/", + "https", + "proxy.example.com", + "8443", + "https://proxy.example.com:8443", + ), + # Case 6: Only X-Forwarded-Host without proto - use original scheme + ( + "http://localhost:4000/", + None, + "proxy.example.com", + None, + "http://proxy.example.com", + ), + # Case 7: Only X-Forwarded-Port without host - preserves original port if present + # (This is safer behavior - X-Forwarded-Port alone is unusual) + ( + "http://localhost:4000/", + None, + None, + "8443", + "http://localhost:4000", # Original port preserved when already present + ), + # Case 8: Complex internal URL with path (path is preserved) + ( + "http://localhost:8888/github/mcp", + "https", + "proxy.example.com", + None, + "https://proxy.example.com/github/mcp", + ), + # Case 9: IPv6 address in X-Forwarded-Host (should not treat :: as port separator) + ( + "http://localhost:4000/", + "https", + "[2001:db8::1]", + None, + "https://[2001:db8::1]", + ), + # Case 10: IPv6 address with port + ( + "http://localhost:4000/", + "https", + "[2001:db8::1]:8080", + None, + "https://[2001:db8::1]:8080", + ), + # Case 11: X-Forwarded-Host already has port, X-Forwarded-Port also provided (host wins) + ( + "http://localhost:4000/", + "https", + "proxy.example.com:9000", + "8443", + "https://proxy.example.com:9000", + ), + # Case 12: Standard proxy setup (most common case) + ( + "http://127.0.0.1:8888/", + "https", + "chatproxy.company.com", + None, + "https://chatproxy.company.com", + ), + # Case 13: Internal URL already has port, X-Forwarded-Port does NOT override + # (safer behavior - preserves original port when X-Forwarded-Host not provided) + ( + "http://localhost:4000/", + None, + None, + "443", + "http://localhost:4000", # Original port preserved + ), + # Case 14: Original URL with existing port in netloc, X-Forwarded-Host replaces it + ( + "http://internal.local:8888/", + "https", + "external.com", + None, + "https://external.com", + ), + ], +) +def test_get_request_base_url_comprehensive( + base_url, x_forwarded_proto, x_forwarded_host, x_forwarded_port, expected_url +): + """Comprehensive test for get_request_base_url with various header combinations""" + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + get_request_base_url, + ) + from fastapi import Request + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + # Create mock request + mock_request = MagicMock(spec=Request) + mock_request.base_url = base_url + + # Build headers dict + headers = {} + if x_forwarded_proto: + headers["X-Forwarded-Proto"] = x_forwarded_proto + if x_forwarded_host: + headers["X-Forwarded-Host"] = x_forwarded_host + if x_forwarded_port: + headers["X-Forwarded-Port"] = x_forwarded_port + + # Mock headers.get() to return our test values + def mock_get(header_name, default=None): + return headers.get(header_name, default) + + mock_request.headers.get = mock_get + + # Test the function + result = get_request_base_url(mock_request) + + # Verify result + assert result == expected_url, ( + f"Expected '{expected_url}' but got '{result}'\n" + f"Input: base_url={base_url}, " + f"X-Forwarded-Proto={x_forwarded_proto}, " + f"X-Forwarded-Host={x_forwarded_host}, " + f"X-Forwarded-Port={x_forwarded_port}" + )