get pass through include config defined pass through

This commit is contained in:
yuneng-jiang 2026-02-10 14:55:37 -08:00
parent 9f8878ee17
commit fc0563fab3
3 changed files with 160 additions and 7 deletions

View file

@ -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):

View file

@ -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())

View file

@ -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():
"""