mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-13 23:11:40 +00:00
fix(mcp): expire temporary OAuth discovery results
This commit is contained in:
parent
bb9b4c4aef
commit
8d5a675878
2 changed files with 47 additions and 0 deletions
|
|
@ -254,6 +254,7 @@ _TRUE_ENV_VALUES: Final = frozenset(("1", "true", "yes", "on"))
|
|||
_OAUTH_DISCOVERY_RETRY_DELAYS_SECONDS: Final = (0.05, 0.15)
|
||||
_OAUTH_DISCOVERY_RETRY_BASE_SECONDS: Final = 30.0
|
||||
_OAUTH_DISCOVERY_RETRY_MAX_SECONDS: Final = 900.0
|
||||
_OAUTH_TEMPORARY_DISCOVERY_TTL_SECONDS: Final = 300.0
|
||||
|
||||
|
||||
def _oauth_discovery_now() -> float:
|
||||
|
|
@ -1898,6 +1899,10 @@ class MCPServerManager:
|
|||
slot: Final = self._oauth_discovery_slot(server_id)
|
||||
return slot is not None and slot.generation == generation
|
||||
|
||||
def _expire_temporary_oauth_discovery(self, server_id: str, generation: int) -> None:
|
||||
if self._oauth_discovery_slot_is_current(server_id, generation):
|
||||
self._remove_oauth_discovery_slot(server_id)
|
||||
|
||||
def _publish_resolved_oauth_server(
|
||||
self,
|
||||
server: MCPServer,
|
||||
|
|
@ -1910,6 +1915,12 @@ class MCPServerManager:
|
|||
elif server.server_id in self.config_mcp_servers:
|
||||
self.config_mcp_servers[server.server_id] = server
|
||||
else:
|
||||
asyncio.get_running_loop().call_later(
|
||||
_OAUTH_TEMPORARY_DISCOVERY_TTL_SECONDS,
|
||||
self._expire_temporary_oauth_discovery,
|
||||
server.server_id,
|
||||
generation,
|
||||
)
|
||||
return server
|
||||
self._remove_oauth_discovery_slot(server.server_id)
|
||||
return server
|
||||
|
|
|
|||
|
|
@ -12758,3 +12758,39 @@ def test_stale_discovery_cannot_overwrite_new_registered_server() -> None:
|
|||
manager._set_oauth_discovery_deferred(original.server_id, True)
|
||||
assert manager._publish_resolved_oauth_server(original, original_slot.generation) is None
|
||||
assert manager.registry[original.server_id] is replacement
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_temporary_oauth_discovery_expires_without_more_requests() -> None:
|
||||
manager: Final = MCPServerManager()
|
||||
server: Final = MCPServer(
|
||||
server_id="expiring-session", name="temporary", url="https://idp.example.com/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.true_passthrough,
|
||||
authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token",
|
||||
)
|
||||
manager._set_oauth_discovery_deferred(server.server_id, True)
|
||||
resolved: Final = await manager.ensure_oauth_metadata_discovered(server)
|
||||
assert manager._oauth_discovery_slot(server.server_id) is not None
|
||||
loop: Final = asyncio.get_running_loop()
|
||||
expired: Final = loop.create_future()
|
||||
with patch.object(loop, "time", return_value=loop.time() + 301):
|
||||
loop.call_later(0, expired.set_result, None)
|
||||
await expired
|
||||
assert resolved.authorization_url == server.authorization_url
|
||||
assert manager._oauth_discovery_slot(server.server_id) is None
|
||||
|
||||
|
||||
def test_old_temporary_discovery_expiry_preserves_replacement() -> None:
|
||||
manager: Final = MCPServerManager()
|
||||
manager._set_oauth_discovery_deferred("reused-session", True)
|
||||
old_slot: Final = manager._oauth_discovery_slot("reused-session")
|
||||
assert old_slot is not None
|
||||
manager._set_oauth_discovery_deferred("reused-session", True)
|
||||
replacement: Final = manager._oauth_discovery_slot("reused-session")
|
||||
manager._expire_temporary_oauth_discovery("reused-session", old_slot.generation)
|
||||
assert manager._oauth_discovery_slot("reused-session") is replacement
|
||||
assert replacement is not None
|
||||
manager._expire_temporary_oauth_discovery("reused-session", replacement.generation)
|
||||
assert manager._oauth_discovery_slot("reused-session") is None
|
||||
manager._expire_temporary_oauth_discovery("reused-session", replacement.generation)
|
||||
assert manager._oauth_discovery_slot("reused-session") is None
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue