mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
refactor(mcp): narrow catalog consistency change to ticket scope
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
b0f1eb656c
commit
51892184a3
7 changed files with 25 additions and 356 deletions
|
|
@ -1,97 +0,0 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
import os
|
||||
import time
|
||||
from collections.abc import Callable, Iterator
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from typing import Final
|
||||
|
||||
from fastapi import HTTPException, Request
|
||||
|
||||
from litellm.proxy.auth.login_throttle import source_group
|
||||
from litellm.proxy.auth.network import TrustedProxyConfig, resolve_client_ip
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class _Window:
|
||||
started: float
|
||||
requests: int = 0
|
||||
active: int = 0
|
||||
|
||||
|
||||
def _limit(name: str, default: int) -> int:
|
||||
value: Final = int(os.environ.get(name, str(default)))
|
||||
if value < 1:
|
||||
raise ValueError(f"{name} must be positive")
|
||||
return value
|
||||
|
||||
|
||||
class MCPAdmissionLimiter:
|
||||
def __init__(self, clock: Callable[[], float] = time.monotonic) -> None:
|
||||
self._clock = clock
|
||||
self._client_rpm = _limit("LITELLM_MCP_PUBLIC_RPM", 120)
|
||||
self._worker_rpm = _limit("LITELLM_MCP_PUBLIC_WORKER_RPM", 600)
|
||||
self._client_active = _limit("LITELLM_MCP_PUBLIC_MAX_IN_FLIGHT", 64)
|
||||
self._worker_active = _limit("LITELLM_MCP_PUBLIC_WORKER_MAX_IN_FLIGHT", 128)
|
||||
self._max_sources = _limit("LITELLM_MCP_PUBLIC_MAX_SOURCES", 4096)
|
||||
self._sources: dict[str, _Window] = {}
|
||||
self._worker = _Window(clock())
|
||||
|
||||
@staticmethod
|
||||
def _reject(retry_after: int) -> None:
|
||||
raise HTTPException(
|
||||
status_code=429,
|
||||
detail="MCP configuration request limit exceeded; retry later",
|
||||
headers={"Retry-After": str(max(1, retry_after))},
|
||||
)
|
||||
|
||||
@contextmanager
|
||||
def admit(self, source: str) -> Iterator[None]:
|
||||
now: Final = self._clock()
|
||||
if now - self._worker.started >= 60:
|
||||
self._worker = _Window(now, active=self._worker.active)
|
||||
self._sources = {
|
||||
key: value for key, value in self._sources.items() if value.active or now - value.started < 60
|
||||
}
|
||||
if self._worker.active >= self._worker_active:
|
||||
self._reject(1)
|
||||
if self._worker.requests >= self._worker_rpm:
|
||||
self._reject(math.ceil(60 - (now - self._worker.started)))
|
||||
previous: Final = self._sources.get(source)
|
||||
if previous is None and len(self._sources) >= self._max_sources:
|
||||
self._reject(math.ceil(60 - (now - self._worker.started)))
|
||||
client: Final = (
|
||||
previous
|
||||
if previous is not None and now - previous.started < 60
|
||||
else _Window(now, active=previous.active if previous is not None else 0)
|
||||
)
|
||||
if client.active >= self._client_active:
|
||||
self._reject(1)
|
||||
if client.requests >= self._client_rpm:
|
||||
self._reject(math.ceil(60 - (now - client.started)))
|
||||
self._sources[source] = client
|
||||
client.requests += 1
|
||||
client.active += 1
|
||||
self._worker.requests += 1
|
||||
self._worker.active += 1
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
self._sources[source].active -= 1
|
||||
self._worker.active -= 1
|
||||
|
||||
|
||||
def admission_source(request: Request) -> str:
|
||||
from litellm.proxy.proxy_server import general_settings # noqa: PLC0415 # runtime proxy configuration
|
||||
|
||||
settings: Final = general_settings or {}
|
||||
config: Final = TrustedProxyConfig.model_validate(
|
||||
{
|
||||
"use_forwarded_for": settings.get("use_x_forwarded_for", False),
|
||||
"trusted_proxy_cidrs": settings.get("mcp_trusted_proxy_ranges") or (),
|
||||
}
|
||||
)
|
||||
client_ip, _ = resolve_client_ip(request, config)
|
||||
return source_group(client_ip or "unknown")
|
||||
|
|
@ -446,7 +446,7 @@ class MCPRequestHandler:
|
|||
Raises:
|
||||
HTTPException: If headers are invalid or missing required headers
|
||||
"""
|
||||
async with global_manager().catalog.operation(request=Request(scope)):
|
||||
async with global_manager().catalog.operation():
|
||||
headers: Final = MCPRequestHandler._safe_get_headers_from_scope(scope)
|
||||
|
||||
# Check if there is an explicit LiteLLM API key (primary header)
|
||||
|
|
|
|||
|
|
@ -10,19 +10,14 @@ from contextlib import asynccontextmanager
|
|||
from contextvars import ContextVar
|
||||
from dataclasses import dataclass, replace
|
||||
from functools import wraps
|
||||
from inspect import signature
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, ParamSpec, TypeVar
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from mcp.types import Tool as SDKTool
|
||||
from pydantic import BaseModel
|
||||
from starlette.requests import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.auth.admission import MCPAdmissionLimiter
|
||||
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing, ServerOutcome
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
from litellm.types.mcp_server.tool_registry import MCPTool
|
||||
|
|
@ -86,7 +81,6 @@ class TargetCatalog:
|
|||
self._arrival_ticket = 0
|
||||
self._completed_ticket = 0
|
||||
self._shared_snapshot: CatalogSnapshot | None = None
|
||||
self._admission: MCPAdmissionLimiter | None = None
|
||||
self._warned_shadowed_config_server_ids: frozenset[str] = frozenset()
|
||||
self._warned_capturing_config_server_ids: frozenset[str] = frozenset()
|
||||
self._operation: ContextVar[tuple[CatalogSnapshot, asyncio.Event, int] | None] = ContextVar(
|
||||
|
|
@ -136,18 +130,8 @@ class TargetCatalog:
|
|||
raise HTTPException(status_code=503, detail="MCP server configuration could not be refreshed")
|
||||
return self._shared_snapshot
|
||||
|
||||
async def _acquire_snapshot(self, request: Request | None) -> CatalogSnapshot:
|
||||
if request is None:
|
||||
return await self._fresh_snapshot()
|
||||
from litellm.proxy._experimental.mcp_server.auth.admission import MCPAdmissionLimiter, admission_source
|
||||
|
||||
if self._admission is None:
|
||||
self._admission = MCPAdmissionLimiter()
|
||||
with self._admission.admit(admission_source(request)):
|
||||
return await self._fresh_snapshot()
|
||||
|
||||
async def list(self, *, request: Request | None = None) -> Mapping[str, MCPServer]:
|
||||
async with self.operation(request=request) as snapshot:
|
||||
async def list(self) -> Mapping[str, MCPServer]:
|
||||
async with self.operation() as snapshot:
|
||||
return snapshot.servers
|
||||
|
||||
def assert_current(self, server: MCPServer) -> None:
|
||||
|
|
@ -168,7 +152,7 @@ class TargetCatalog:
|
|||
raise HTTPException(status_code=503, detail="MCP server configuration changed; retry the operation")
|
||||
|
||||
@asynccontextmanager
|
||||
async def operation(self, *, request: Request | None = None) -> AsyncIterator[CatalogSnapshot]:
|
||||
async def operation(self) -> AsyncIterator[CatalogSnapshot]:
|
||||
current: Final = self.current()
|
||||
scoped: Final = self._operation.get()
|
||||
if current is not None and scoped is not None and scoped[2] == id(asyncio.current_task()):
|
||||
|
|
@ -176,7 +160,7 @@ class TargetCatalog:
|
|||
return
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
|
||||
|
||||
shared: Final = await self._acquire_snapshot(request)
|
||||
shared: Final = await self._fresh_snapshot()
|
||||
snapshot: Final = replace(
|
||||
shared,
|
||||
servers=MappingProxyType({key: value.model_copy(deep=True) for key, value in shared.servers.items()}),
|
||||
|
|
@ -242,27 +226,6 @@ class TargetCatalog:
|
|||
_check_oauth_revision(selected, resolved)
|
||||
return resolved
|
||||
|
||||
@staticmethod
|
||||
async def aggregate_list(
|
||||
servers: Sequence[MCPServer],
|
||||
fetch: Callable[[MCPServer], Awaitable[tuple[list[SDKTool], ServerOutcome]]],
|
||||
server_key: Callable[[MCPServer], str],
|
||||
) -> AggregateToolListing:
|
||||
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing
|
||||
|
||||
tasks: Final = tuple(asyncio.ensure_future(fetch(server)) for server in servers)
|
||||
try:
|
||||
results: Final = await asyncio.gather(*tasks)
|
||||
return AggregateToolListing(
|
||||
tools=[tool for tools, _ in results for tool in tools],
|
||||
outcomes={server_key(server): outcome for server, (_, outcome) in zip(servers, results)},
|
||||
)
|
||||
finally:
|
||||
for task in tasks:
|
||||
if not task.done():
|
||||
task.cancel()
|
||||
await asyncio.gather(*tasks, return_exceptions=True)
|
||||
|
||||
async def reload(self) -> None:
|
||||
async with self._refresh_lock:
|
||||
await self._publish_refresh()
|
||||
|
|
@ -559,18 +522,6 @@ def global_manager() -> MCPServerManager:
|
|||
return global_mcp_server_manager
|
||||
|
||||
|
||||
def public_catalog_operation(function: Callable[_P, Awaitable[_R]]) -> Callable[_P, Awaitable[_R]]:
|
||||
parameters: Final = signature(function)
|
||||
|
||||
@wraps(function)
|
||||
async def wrapped(*args: _P.args, **kwargs: _P.kwargs) -> _R: # kwargs-ok: preserves ParamSpec
|
||||
from starlette.requests import Request
|
||||
|
||||
arguments: Final = parameters.bind(*args, **kwargs).arguments
|
||||
request: Final = arguments.get("request")
|
||||
if not isinstance(request, Request):
|
||||
raise TypeError("Public MCP operations require a Request")
|
||||
async with global_manager().catalog.operation(request=request):
|
||||
return await function(*args, **kwargs)
|
||||
|
||||
return wrapped
|
||||
public_catalog_operation: Final[Callable[[Callable[_P, Awaitable[_R]]], Callable[_P, Awaitable[_R]]]] = (
|
||||
catalog_operation(global_manager)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -2919,7 +2919,7 @@ async def register_client(request: Request, mcp_server_name: str | None = None):
|
|||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
async with global_mcp_server_manager.catalog.operation(request=request):
|
||||
async with global_mcp_server_manager.catalog.operation():
|
||||
dummy_return: Final = {
|
||||
"client_id": mcp_server_name or "dummy_client",
|
||||
"client_secret": "dummy",
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
"""Shared MCP operation policy and dispatch."""
|
||||
|
||||
import asyncio
|
||||
import traceback
|
||||
import types
|
||||
import uuid
|
||||
|
|
@ -49,7 +50,7 @@ from litellm.proxy._experimental.mcp_server.byok_credential_cache import (
|
|||
cache_byok_credential,
|
||||
get_cached_byok_credential,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.catalog import TargetCatalog, catalog_operation, global_manager
|
||||
from litellm.proxy._experimental.mcp_server.catalog import catalog_operation, global_manager
|
||||
from litellm.proxy._experimental.mcp_server.contracts import (
|
||||
AuthorizedToolCall,
|
||||
OperationContext,
|
||||
|
|
@ -1153,17 +1154,24 @@ async def _get_tools_from_mcp_servers(
|
|||
verbose_logger.exception("Error getting tools from server %s: %s", server.name, e)
|
||||
return [], classify_list_exception(e)
|
||||
|
||||
listing: Final = await TargetCatalog.aggregate_list(
|
||||
allowed_mcp_servers, _fetch_and_filter_server_tools, _aggregate_server_key
|
||||
)
|
||||
all_tools: Final = listing.tools
|
||||
server_outcomes: Final = listing.outcomes
|
||||
# Fetch tools from all servers in parallel
|
||||
tasks: Final = [_fetch_and_filter_server_tools(server) for server in allowed_mcp_servers]
|
||||
results: Final = await asyncio.gather(*tasks)
|
||||
|
||||
# Flatten results into single list
|
||||
all_tools: Final[list[MCPTool]] = [tool for tools, _ in results for tool in tools]
|
||||
server_outcomes: Final[dict[str, ServerOutcome]] = {
|
||||
_aggregate_server_key(server): outcome
|
||||
for server, (_, outcome) in zip(allowed_mcp_servers, results)
|
||||
if server is not None
|
||||
}
|
||||
|
||||
# If logging is enabled, enrich spend_logs_metadata with counts
|
||||
if litellm_logging_obj:
|
||||
per_server_tool_counts: Final[dict[str, int]] = {
|
||||
key: outcome.tool_count if isinstance(outcome, ServerListOk) else 0
|
||||
for key, outcome in server_outcomes.items()
|
||||
_aggregate_server_key(server): len(server_tools)
|
||||
for server, (server_tools, _) in zip(allowed_mcp_servers, results)
|
||||
if server is not None
|
||||
}
|
||||
|
||||
metadata_dict: Final = litellm_logging_obj.model_call_details.get("metadata")
|
||||
|
|
|
|||
|
|
@ -1,89 +0,0 @@
|
|||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException, Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.auth.admission import MCPAdmissionLimiter, admission_source
|
||||
|
||||
|
||||
def test_client_and_worker_rate_budgets_recover_after_window(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_MCP_PUBLIC_RPM", "2")
|
||||
monkeypatch.setenv("LITELLM_MCP_PUBLIC_WORKER_RPM", "3")
|
||||
clock = Mock(return_value=0.0)
|
||||
limiter = MCPAdmissionLimiter(clock)
|
||||
for source in ("a", "a"):
|
||||
with limiter.admit(source):
|
||||
pass
|
||||
with pytest.raises(HTTPException) as client_error, limiter.admit("a"):
|
||||
pytest.fail("client budget must reject")
|
||||
assert client_error.value.status_code == 429
|
||||
assert client_error.value.headers == {"Retry-After": "60"}
|
||||
with limiter.admit("b"):
|
||||
pass
|
||||
clock.return_value = 59.1
|
||||
with pytest.raises(HTTPException) as worker_error, limiter.admit("c"):
|
||||
pytest.fail("worker budget must reject rotated sources")
|
||||
assert worker_error.value.headers == {"Retry-After": "1"}
|
||||
clock.return_value = 60.0
|
||||
with limiter.admit("a"), limiter.admit("c"):
|
||||
pass
|
||||
|
||||
|
||||
def test_inflight_limits_release_on_failure_across_window_boundary(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_MCP_PUBLIC_MAX_IN_FLIGHT", "1")
|
||||
monkeypatch.setenv("LITELLM_MCP_PUBLIC_WORKER_MAX_IN_FLIGHT", "2")
|
||||
clock = Mock(return_value=0.0)
|
||||
limiter = MCPAdmissionLimiter(clock)
|
||||
with limiter.admit("a"):
|
||||
with pytest.raises(HTTPException) as client_error, limiter.admit("a"):
|
||||
pytest.fail("one client cannot occupy another permit")
|
||||
assert client_error.value.headers == {"Retry-After": "1"}
|
||||
with limiter.admit("b"):
|
||||
clock.return_value = 60.0
|
||||
with pytest.raises(HTTPException) as worker_error, limiter.admit("c"):
|
||||
pytest.fail("rotating sources cannot exceed active work")
|
||||
assert worker_error.value.status_code == 429
|
||||
with pytest.raises(RuntimeError, match="upstream"), limiter.admit("c"):
|
||||
raise RuntimeError("upstream")
|
||||
with limiter.admit("a"), limiter.admit("c"):
|
||||
pass
|
||||
|
||||
|
||||
def test_source_capacity_never_evicts_live_budget(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_MCP_PUBLIC_MAX_SOURCES", "2")
|
||||
clock = Mock(return_value=0.0)
|
||||
limiter = MCPAdmissionLimiter(clock)
|
||||
with limiter.admit("a"):
|
||||
with limiter.admit("b"):
|
||||
pass
|
||||
for index in range(100):
|
||||
with pytest.raises(HTTPException) as error, limiter.admit(str(index)):
|
||||
pytest.fail("source spray must be bounded")
|
||||
assert error.value.status_code == 429
|
||||
clock.return_value = 60.0
|
||||
with limiter.admit("c"):
|
||||
pass
|
||||
with limiter.admit("a"):
|
||||
pass
|
||||
|
||||
|
||||
@pytest.mark.parametrize("value", ["0", "-1", "invalid"])
|
||||
def test_invalid_limits_are_not_silently_disabled(monkeypatch, value):
|
||||
monkeypatch.setenv("LITELLM_MCP_PUBLIC_RPM", value)
|
||||
with pytest.raises(ValueError, match=r"must be positive|invalid literal"):
|
||||
MCPAdmissionLimiter()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("peer,trusted,xff,expected", [
|
||||
("198.51.100.1", [], "203.0.113.1", "198.51.100.1"),
|
||||
("10.0.0.1", ["10.0.0.0/8"], "203.0.113.1, 198.51.100.1", "198.51.100.1"),
|
||||
("2001:db8::1", [], "", "2001:db8::/64"),
|
||||
(None, [], "203.0.113.1", "unknown"),
|
||||
])
|
||||
def test_source_identity_ignores_untrusted_forwarded_addresses(monkeypatch, peer, trusted, xff, expected):
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
monkeypatch.setattr(proxy_server, "general_settings", {"use_x_forwarded_for": True, "mcp_trusted_proxy_ranges": trusted})
|
||||
request = Request({"type": "http", "client": (peer, 1234) if peer else None,
|
||||
"headers": [(b"x-forwarded-for", xff.encode())]})
|
||||
assert admission_source(request) == expected
|
||||
|
|
@ -15489,55 +15489,6 @@ async def test_catalog_fresh_lookup_does_not_fall_back_to_stale_grants_when_data
|
|||
assert manager.registry == {server.server_id: server}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_catalog_list_failure_cancels_and_joins_other_upstream_fetches():
|
||||
from litellm.proxy._experimental.mcp_server.catalog import TargetCatalog
|
||||
|
||||
started: Final = asyncio.Event()
|
||||
stopped: Final = asyncio.Event()
|
||||
pending: Final = MCPServer(server_id="pending", name="pending", transport=MCPTransport.http)
|
||||
broken: Final = MCPServer(server_id="broken", name="broken", transport=MCPTransport.http)
|
||||
|
||||
async def fetch(server: MCPServer):
|
||||
if server.server_id == broken.server_id:
|
||||
await started.wait()
|
||||
raise RuntimeError("unexpected fetch failure")
|
||||
started.set()
|
||||
try:
|
||||
await asyncio.Event().wait()
|
||||
finally:
|
||||
stopped.set()
|
||||
|
||||
with pytest.raises(RuntimeError, match="unexpected fetch failure"):
|
||||
await TargetCatalog.aggregate_list((pending, broken), fetch, lambda server: server.server_id)
|
||||
assert stopped.is_set()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("headers", [[], [(b"authorization", b"Bearer arbitrary")]])
|
||||
async def test_public_catalog_budget_rejects_before_database_read(monkeypatch, headers):
|
||||
from starlette.requests import Request
|
||||
from litellm.proxy._experimental.mcp_server import mcp_server_manager
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import oauth_authorization_server_mcp
|
||||
|
||||
monkeypatch.setenv("LITELLM_MCP_PUBLIC_RPM", "2")
|
||||
read_rows = AsyncMock(return_value=[])
|
||||
_catalog_database(monkeypatch, read_rows)
|
||||
manager = MCPServerManager()
|
||||
monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager)
|
||||
request = Request({"type": "http", "method": "GET", "path": "/.well-known/oauth-authorization-server/missing",
|
||||
"headers": headers, "client": ("198.51.100.1", 1234), "scheme": "https",
|
||||
"server": ("gateway.example", 443), "query_string": b""})
|
||||
for expected in (404, 404, 429):
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await oauth_authorization_server_mcp(request=request, mcp_server_name="missing")
|
||||
assert exc.value.status_code == expected
|
||||
assert int(exc.value.headers["Retry-After"]) > 0
|
||||
assert read_rows.await_count == 2
|
||||
await manager.catalog.list()
|
||||
assert read_rows.await_count == 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_catalog_cancelled_registration_does_not_publish_partial_handlers(monkeypatch):
|
||||
from litellm.proxy._experimental.mcp_server import tool_registry
|
||||
|
|
@ -15558,58 +15509,3 @@ async def test_catalog_cancelled_registration_does_not_publish_partial_handlers(
|
|||
assert [tool.name for tool in registry.list_tools()] == ["existing"]
|
||||
assert manager.registry == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("cancel", [False, True])
|
||||
async def test_public_catalog_inflight_rejection_precedes_read_and_releases_permit(monkeypatch, cancel):
|
||||
from starlette.requests import Request
|
||||
|
||||
monkeypatch.setenv("LITELLM_MCP_PUBLIC_WORKER_MAX_IN_FLIGHT", "1")
|
||||
started = asyncio.Event()
|
||||
release = asyncio.Event()
|
||||
|
||||
async def blocked_rows(**kwargs):
|
||||
started.set()
|
||||
await release.wait()
|
||||
return [_catalog_row()]
|
||||
|
||||
read_rows = AsyncMock(side_effect=blocked_rows)
|
||||
_catalog_database(monkeypatch, read_rows)
|
||||
manager = MCPServerManager()
|
||||
request = Request({"type": "http", "client": ("198.51.100.1", 1234), "headers": []})
|
||||
first = asyncio.create_task(manager.catalog.list(request=request))
|
||||
await asyncio.wait_for(started.wait(), 1)
|
||||
with pytest.raises(HTTPException) as error:
|
||||
await manager.catalog.list(request=request)
|
||||
assert error.value.status_code == 429
|
||||
assert error.value.headers == {"Retry-After": "1"}
|
||||
assert read_rows.await_count == 1
|
||||
authenticated = asyncio.create_task(manager.catalog.list())
|
||||
if cancel:
|
||||
first.cancel()
|
||||
release.set()
|
||||
if cancel:
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await first
|
||||
else:
|
||||
assert "catalog-server" in await first
|
||||
assert "catalog-server" in await asyncio.wait_for(authenticated, 1)
|
||||
assert "catalog-server" in await manager.catalog.list(request=request)
|
||||
assert read_rows.await_count == 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_public_catalog_releases_permit_before_operation_and_nested_use(monkeypatch):
|
||||
from starlette.requests import Request
|
||||
|
||||
monkeypatch.setenv("LITELLM_MCP_PUBLIC_WORKER_MAX_IN_FLIGHT", "1")
|
||||
read_rows = AsyncMock(return_value=[_catalog_row()])
|
||||
_catalog_database(monkeypatch, read_rows)
|
||||
manager = MCPServerManager()
|
||||
request = Request({"type": "http", "client": ("198.51.100.1", 1234), "headers": []})
|
||||
async with manager.catalog.operation(request=request):
|
||||
async with manager.catalog.operation(request=request):
|
||||
assert read_rows.await_count == 1
|
||||
other = await asyncio.create_task(manager.catalog.list(request=request))
|
||||
assert "catalog-server" in other
|
||||
assert read_rows.await_count == 2
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue