mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
get pass through include config defined pass through
This commit is contained in:
parent
9f8878ee17
commit
fc0563fab3
3 changed files with 160 additions and 7 deletions
|
|
@ -1910,6 +1910,10 @@ class PassThroughGenericEndpoint(LiteLLMPydanticObjectBase):
|
|||
default=None,
|
||||
description="Guardrails configuration for this passthrough endpoint. Dict keys are guardrail names, values are optional settings for field targeting. When set, all org/team/key level guardrails will also execute. Defaults to None (no guardrails execute).",
|
||||
)
|
||||
is_from_config: bool = Field(
|
||||
default=False,
|
||||
description="True if this endpoint is defined in the config file, False if from DB. Config-defined endpoints cannot be edited via the UI.",
|
||||
)
|
||||
|
||||
|
||||
class PassThroughEndpointResponse(LiteLLMPydanticObjectBase):
|
||||
|
|
|
|||
|
|
@ -2201,6 +2201,31 @@ async def initialize_pass_through_endpoints(
|
|||
InitPassThroughEndpointHelpers.remove_endpoint_routes(endpoint_key)
|
||||
|
||||
|
||||
def _get_pass_through_endpoints_from_config() -> List[PassThroughGenericEndpoint]:
|
||||
"""
|
||||
Get pass-through endpoints defined in the config file.
|
||||
These are read-only and cannot be edited via the UI.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import config_passthrough_endpoints
|
||||
|
||||
if config_passthrough_endpoints is None or len(config_passthrough_endpoints) == 0:
|
||||
return []
|
||||
|
||||
returned_endpoints: List[PassThroughGenericEndpoint] = []
|
||||
for endpoint in config_passthrough_endpoints:
|
||||
if isinstance(endpoint, dict):
|
||||
endpoint_dict = dict(endpoint)
|
||||
endpoint_dict["is_from_config"] = True
|
||||
returned_endpoints.append(PassThroughGenericEndpoint(**endpoint_dict))
|
||||
elif isinstance(endpoint, PassThroughGenericEndpoint):
|
||||
# Create a copy with is_from_config=True
|
||||
endpoint_dict = endpoint.model_dump()
|
||||
endpoint_dict["is_from_config"] = True
|
||||
returned_endpoints.append(PassThroughGenericEndpoint(**endpoint_dict))
|
||||
|
||||
return returned_endpoints
|
||||
|
||||
|
||||
async def _get_pass_through_endpoints_from_db(
|
||||
endpoint_id: Optional[str] = None,
|
||||
user_api_key_dict: Optional[UserAPIKeyAuth] = None,
|
||||
|
|
@ -2223,17 +2248,27 @@ async def _get_pass_through_endpoints_from_db(
|
|||
|
||||
returned_endpoints: List[PassThroughGenericEndpoint] = []
|
||||
if endpoint_id is None:
|
||||
# Return all endpoints
|
||||
# Return all endpoints from DB, mark as not from config
|
||||
for endpoint in pass_through_endpoint_data:
|
||||
if isinstance(endpoint, dict):
|
||||
returned_endpoints.append(PassThroughGenericEndpoint(**endpoint))
|
||||
endpoint_dict = dict(endpoint)
|
||||
endpoint_dict["is_from_config"] = False
|
||||
returned_endpoints.append(PassThroughGenericEndpoint(**endpoint_dict))
|
||||
elif isinstance(endpoint, PassThroughGenericEndpoint):
|
||||
returned_endpoints.append(endpoint)
|
||||
endpoint_dict = endpoint.model_dump()
|
||||
endpoint_dict["is_from_config"] = False
|
||||
returned_endpoints.append(PassThroughGenericEndpoint(**endpoint_dict))
|
||||
else:
|
||||
# Find specific endpoint by ID
|
||||
found_endpoint = _find_endpoint_by_id(pass_through_endpoint_data, endpoint_id)
|
||||
if found_endpoint is not None:
|
||||
returned_endpoints.append(found_endpoint)
|
||||
endpoint_dict = (
|
||||
found_endpoint.model_dump()
|
||||
if isinstance(found_endpoint, PassThroughGenericEndpoint)
|
||||
else dict(found_endpoint)
|
||||
)
|
||||
endpoint_dict["is_from_config"] = False
|
||||
returned_endpoints.append(PassThroughGenericEndpoint(**endpoint_dict))
|
||||
|
||||
return returned_endpoints
|
||||
|
||||
|
|
@ -2312,10 +2347,25 @@ async def get_pass_through_endpoints(
|
|||
detail={"error": CommonProxyErrors.db_not_connected_error.value},
|
||||
)
|
||||
|
||||
pass_through_endpoints = await _get_pass_through_endpoints_from_db(
|
||||
# Get endpoints from DB (editable via UI)
|
||||
db_endpoints = await _get_pass_through_endpoints_from_db(
|
||||
endpoint_id=endpoint_id, user_api_key_dict=user_api_key_dict
|
||||
)
|
||||
|
||||
# Get endpoints from config file (read-only, not editable via UI)
|
||||
config_endpoints = _get_pass_through_endpoints_from_config()
|
||||
|
||||
# Merge: config endpoints not in DB + all DB endpoints (DB overrides config for same path)
|
||||
db_paths = {ep.path for ep in db_endpoints}
|
||||
config_only_endpoints = [
|
||||
ep for ep in config_endpoints if ep.path not in db_paths
|
||||
]
|
||||
if endpoint_id is not None:
|
||||
# When filtering by endpoint_id, only return if found in DB (config endpoints use generated IDs)
|
||||
pass_through_endpoints = db_endpoints
|
||||
else:
|
||||
pass_through_endpoints = config_only_endpoints + db_endpoints
|
||||
|
||||
if team_id is not None:
|
||||
pass_through_endpoints = await _filter_endpoints_by_team_allowed_routes(
|
||||
team_id=team_id,
|
||||
|
|
@ -2392,7 +2442,8 @@ async def update_pass_through_endpoints(
|
|||
)
|
||||
|
||||
# Get the update data as dict, excluding None values for partial updates
|
||||
update_data = data.model_dump(exclude_none=True)
|
||||
# Exclude is_from_config as it's a response-only field (computed at read time)
|
||||
update_data = data.model_dump(exclude_none=True, exclude={"is_from_config"})
|
||||
|
||||
# Start with existing endpoint data
|
||||
endpoint_dict = found_endpoint.model_dump()
|
||||
|
|
@ -2404,6 +2455,9 @@ async def update_pass_through_endpoints(
|
|||
if "id" not in update_data and found_endpoint.id is not None:
|
||||
endpoint_dict["id"] = found_endpoint.id
|
||||
|
||||
# Remove is_from_config before saving - it's a response-only field (computed at read time)
|
||||
endpoint_dict.pop("is_from_config", None)
|
||||
|
||||
# Create updated endpoint object
|
||||
updated_endpoint = PassThroughGenericEndpoint(**endpoint_dict)
|
||||
|
||||
|
|
@ -2490,7 +2544,8 @@ async def create_pass_through_endpoints(
|
|||
)
|
||||
|
||||
## Auto-generate ID if not provided
|
||||
data_dict = data.model_dump()
|
||||
# Exclude is_from_config as it's a response-only field (computed at read time)
|
||||
data_dict = data.model_dump(exclude={"is_from_config"})
|
||||
if data_dict.get("id") is None:
|
||||
data_dict["id"] = str(uuid.uuid4())
|
||||
|
||||
|
|
|
|||
|
|
@ -1316,6 +1316,100 @@ async def test_delete_pass_through_endpoint_not_found():
|
|||
assert "not found" in str(exc_info.value.detail).lower()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_pass_through_endpoints_includes_config_and_db():
|
||||
"""
|
||||
Test that get_pass_through_endpoints returns both config-defined and DB endpoints,
|
||||
with correct is_from_config flag. Config-only endpoints have is_from_config=True,
|
||||
DB endpoints have is_from_config=False. When same path exists in both, DB overrides.
|
||||
"""
|
||||
from litellm.proxy._types import (
|
||||
PassThroughEndpointResponse,
|
||||
PassThroughGenericEndpoint,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
|
||||
get_pass_through_endpoints,
|
||||
)
|
||||
|
||||
# Config-defined endpoints (from config file)
|
||||
config_endpoints = [
|
||||
{
|
||||
"path": "/v1/rerank",
|
||||
"target": "https://api.cohere.com/v1/rerank",
|
||||
"headers": {"content-type": "application/json"},
|
||||
},
|
||||
{
|
||||
"path": "/v1/config-only",
|
||||
"target": "https://config.example.com/api",
|
||||
"headers": {},
|
||||
},
|
||||
]
|
||||
|
||||
# DB endpoints (one overlaps with config path, one is DB-only)
|
||||
db_endpoints = [
|
||||
{
|
||||
"id": "db-endpoint-1",
|
||||
"path": "/v1/rerank", # Same as config - DB should override
|
||||
"target": "https://db-override.com/v1/rerank",
|
||||
"headers": {},
|
||||
"include_subpath": False,
|
||||
},
|
||||
{
|
||||
"id": "db-endpoint-2",
|
||||
"path": "/db/only",
|
||||
"target": "https://db-only.example.com/api",
|
||||
"headers": {},
|
||||
"include_subpath": False,
|
||||
},
|
||||
]
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.prisma_client",
|
||||
MagicMock(),
|
||||
):
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints._get_pass_through_endpoints_from_db",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get_db:
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints._get_pass_through_endpoints_from_config"
|
||||
) as mock_get_config:
|
||||
db_objects = [
|
||||
PassThroughGenericEndpoint(**ep, is_from_config=False)
|
||||
for ep in db_endpoints
|
||||
]
|
||||
config_objects = [
|
||||
PassThroughGenericEndpoint(**ep, is_from_config=True)
|
||||
for ep in config_endpoints
|
||||
]
|
||||
mock_get_db.return_value = db_objects
|
||||
mock_get_config.return_value = config_objects
|
||||
|
||||
mock_user = MagicMock(spec=UserAPIKeyAuth)
|
||||
|
||||
result = await get_pass_through_endpoints(
|
||||
endpoint_id=None,
|
||||
user_api_key_dict=mock_user,
|
||||
team_id=None,
|
||||
)
|
||||
|
||||
assert isinstance(result, PassThroughEndpointResponse)
|
||||
# config_only: /v1/config-only (not in db_paths)
|
||||
# db: /v1/rerank (overrides config), /db/only
|
||||
# So we should have: /v1/config-only (from config) + /v1/rerank + /db/only (from db)
|
||||
assert len(result.endpoints) == 3
|
||||
|
||||
# Check is_from_config values
|
||||
by_path = {ep.path: ep for ep in result.endpoints}
|
||||
assert by_path["/v1/config-only"].is_from_config is True
|
||||
assert by_path["/v1/rerank"].is_from_config is False # DB overrides
|
||||
assert by_path["/db/only"].is_from_config is False
|
||||
|
||||
# Verify DB override: /v1/rerank should have DB target
|
||||
assert by_path["/v1/rerank"].target == "https://db-override.com/v1/rerank"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_pass_through_endpoint_empty_list():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue