mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
Improve(mcp): respect X-Forwarded- headers in OAuth endpoints (#16036)
* 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 <noreply@anthropic.com> * 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 <noreply@anthropic.com> * 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 <noreply@anthropic.com>
This commit is contained in:
parent
eb0e4f34dc
commit
5e10ea4136
2 changed files with 629 additions and 12 deletions
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue