mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-17 23:51:30 +00:00
Merge remote-tracking branch 'origin/main' into litellm_model_access_denied_message
This commit is contained in:
commit
b7af51cc4a
9 changed files with 519 additions and 20 deletions
|
|
@ -318,6 +318,8 @@ WEBSOCKET_CLOSE_REASON_MAX_BYTES: Final = 123
|
|||
BEDROCK_REALTIME_PENDING_SESSION_UPDATE_SCOPE_KEY: Final = "litellm.bedrock_realtime.pending_session_update"
|
||||
BEDROCK_REALTIME_SESSION_COMMITTED_SCOPE_KEY: Final = "litellm.bedrock_realtime.session_committed"
|
||||
BEDROCK_REALTIME_COMMITTED_FAILURE_SCOPE_KEY: Final = "litellm.bedrock_realtime.committed_failure"
|
||||
CLIENT_REQUESTED_MODEL_SCOPE_KEY: Final = "litellm.client_requested_model"
|
||||
MODEL_GROUP_ALIAS_RESOLVED_SCOPE_KEY: Final = "litellm.model_group_alias_resolved"
|
||||
REALTIME_SESSION_SUCCESS_LOGGED_KEY: Final = "realtime_session_success_logged"
|
||||
REALTIME_SESSION_FAILURE_LOGGED_KEY: Final = "realtime_session_failure_logged"
|
||||
|
||||
|
|
|
|||
|
|
@ -1967,6 +1967,11 @@ def request_dispatched_to_pass_through_endpoint(request: Request | None) -> bool
|
|||
return getattr(endpoint, LITELLM_PASS_THROUGH_ENDPOINT_MARKER, False) is True
|
||||
|
||||
|
||||
def request_dispatched_to_provider_pass_through(request: Request) -> bool:
|
||||
"""Built-in provider pass-through handlers (``/anthropic/{endpoint:path}``, ...) bind ``endpoint``."""
|
||||
return "endpoint" in request.path_params
|
||||
|
||||
|
||||
def get_model_from_request(
|
||||
request_data: dict,
|
||||
route: str,
|
||||
|
|
|
|||
|
|
@ -26,11 +26,13 @@ from litellm._logging import verbose_logger, verbose_proxy_logger
|
|||
from litellm._service_logger import ServiceLogging
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
from litellm.constants import (
|
||||
CLIENT_REQUESTED_MODEL_SCOPE_KEY,
|
||||
GLOBAL_PROXY_SPEND_CACHE_KEY,
|
||||
INVALID_VIRTUAL_KEY_ERROR_MARKER,
|
||||
INVALID_VIRTUAL_KEY_ERROR_MESSAGE,
|
||||
LITELLM_PROXY_BUDGET_NAME,
|
||||
LITELLM_PROXY_MASTER_KEY_ALIAS,
|
||||
MODEL_GROUP_ALIAS_RESOLVED_SCOPE_KEY,
|
||||
)
|
||||
from litellm.integrations.otel.model.config import is_otel_v2_enabled
|
||||
from litellm.integrations.otel.runtime import phase_span, seed_request_identity
|
||||
|
|
@ -76,6 +78,8 @@ from litellm.proxy.auth.auth_utils import (
|
|||
iter_request_fallback_targets,
|
||||
normalize_request_route,
|
||||
pre_db_read_auth_checks,
|
||||
request_dispatched_to_pass_through_endpoint,
|
||||
request_dispatched_to_provider_pass_through,
|
||||
route_in_additonal_public_routes,
|
||||
)
|
||||
from litellm.proxy.auth.handle_jwt import JWTAuthManager, JWTHandler
|
||||
|
|
@ -103,6 +107,7 @@ from litellm.proxy.common_utils.http_parsing_utils import (
|
|||
_safe_set_request_parsed_body,
|
||||
populate_request_with_path_params,
|
||||
read_raw_json_body,
|
||||
rewrite_request_model,
|
||||
)
|
||||
from litellm.proxy.common_utils.model_listing_utils import claude_code_requested_group
|
||||
from litellm.proxy.common_utils.realtime_utils import _realtime_request_body
|
||||
|
|
@ -124,6 +129,7 @@ from litellm.proxy.utils import (
|
|||
normalize_route_for_root_path,
|
||||
)
|
||||
from litellm.repositories.table_repositories import TeamMembershipRepository
|
||||
from litellm.router_utils.common_utils import resolve_model_group_alias
|
||||
from litellm.secret_managers.main import get_secret_bool
|
||||
from litellm.types.services import ServiceTypes
|
||||
|
||||
|
|
@ -235,11 +241,45 @@ async def _normalize_claude_model(
|
|||
request.scope[_CLAUDE_MODEL_NORMALIZED] = True
|
||||
if source is None:
|
||||
return
|
||||
request_data["model"] = source
|
||||
_safe_set_request_parsed_body(request=request, parsed_body=request_data)
|
||||
if request is not None:
|
||||
request._json = request_data
|
||||
request._body = orjson.dumps(request_data)
|
||||
rewrite_request_model(request_data, request, source)
|
||||
|
||||
|
||||
async def _resolve_router_settings_model_group_alias(
|
||||
request_data: dict[str, object], # mutable-ok: the request body is rewritten in place for every downstream reader
|
||||
valid_token: UserAPIKeyAuth,
|
||||
request: Request | None,
|
||||
route: str,
|
||||
) -> None:
|
||||
"""Rewrite the requested model through the key's or team's ``router_settings.model_group_alias``
|
||||
before the allowlist checks, so they authorize the model group the request is routed to.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import llm_router, prisma_client, proxy_config, proxy_logging_obj
|
||||
|
||||
if request is None or llm_router is None or not RouteChecks.is_llm_api_route(route=route):
|
||||
return
|
||||
if request.scope.get(MODEL_GROUP_ALIAS_RESOLVED_SCOPE_KEY) is True:
|
||||
return
|
||||
request.scope[MODEL_GROUP_ALIAS_RESOLVED_SCOPE_KEY] = True
|
||||
if request_dispatched_to_pass_through_endpoint(request) or request_dispatched_to_provider_pass_through(request):
|
||||
return
|
||||
requested: Final = request_data.get("model")
|
||||
if not isinstance(requested, str) or await read_raw_json_body(request=request) is None:
|
||||
return
|
||||
settings: Final = await proxy_config.get_hierarchical_router_settings(
|
||||
user_api_key_dict=valid_token, prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj
|
||||
)
|
||||
if not isinstance(settings, Mapping):
|
||||
return
|
||||
target: Final = resolve_model_group_alias(settings.get("model_group_alias"), requested)
|
||||
if target is None or target == requested:
|
||||
return
|
||||
verbose_proxy_logger.debug(
|
||||
"router_settings.model_group_alias resolved %s -> %s before auth",
|
||||
requested.replace("\r", "").replace("\n", ""),
|
||||
target.replace("\r", "").replace("\n", ""),
|
||||
)
|
||||
request.scope.setdefault(CLIENT_REQUESTED_MODEL_SCOPE_KEY, requested)
|
||||
rewrite_request_model(request_data, request, target)
|
||||
|
||||
|
||||
def _get_model_names_for_budget_checks(
|
||||
|
|
@ -2990,6 +3030,7 @@ async def _authorize_authenticated_request(
|
|||
## ENSURE DISABLE ROUTE WORKS ACROSS ALL USER AUTH FLOWS ##
|
||||
RouteChecks.should_call_route(route=route, valid_token=user_api_key_auth_obj, request=request)
|
||||
await _normalize_claude_model(request_data, user_api_key_auth_obj, request, route)
|
||||
await _resolve_router_settings_model_group_alias(request_data, user_api_key_auth_obj, request, route)
|
||||
|
||||
# Single authorization point. Builder paths MUST NOT call common_checks.
|
||||
# Route through the same exception handler the builder uses so
|
||||
|
|
@ -3376,6 +3417,7 @@ async def _enforce_key_and_fallback_model_access(
|
|||
Not included in common_checks — common_checks enforces team/user/project model access only.
|
||||
"""
|
||||
await _normalize_claude_model(request_data, valid_token, request, route)
|
||||
await _resolve_router_settings_model_group_alias(request_data, valid_token, request, route)
|
||||
config: Final = valid_token.config
|
||||
|
||||
if config != {}:
|
||||
|
|
|
|||
|
|
@ -56,6 +56,7 @@ from litellm.proxy.common_utils.callback_utils import (
|
|||
get_logging_caching_headers,
|
||||
get_remaining_tokens_and_requests_from_request_data,
|
||||
)
|
||||
from litellm.proxy.common_utils.http_parsing_utils import get_client_requested_model
|
||||
from litellm.proxy.common_utils.openai_error_payload import (
|
||||
attribute_of,
|
||||
error_status_code,
|
||||
|
|
@ -622,9 +623,9 @@ async def _resolve_per_request_model_group_alias(
|
|||
holds the global config map and is shared across requests, so a per-request
|
||||
map has to be applied here instead of being forwarded to the Router.
|
||||
|
||||
Model access was authorized against the requested group, so the target is
|
||||
authorized in its own right before the rewrite; a key that may not call the
|
||||
target gets the usual 403 rather than being quietly served it.
|
||||
Auth already rewrote the body through this map for LLM API routes, so this is
|
||||
a fallback for callers that skipped it; the target is authorized in its own
|
||||
right before the rewrite, so a key that may not call it gets the usual 403.
|
||||
|
||||
Returns the target model group, or None when no alias applies.
|
||||
"""
|
||||
|
|
@ -1451,10 +1452,13 @@ def _has_attribute_error_in_chain(exc: Exception) -> bool:
|
|||
_CLIENT_DISCONNECT_DETAIL: Final = "Client disconnected the request"
|
||||
|
||||
|
||||
def _log_llm_api_exception(e: Exception) -> None:
|
||||
def _log_llm_api_exception(e: Exception, litellm_call_id: str | None) -> None:
|
||||
if getattr(e, "status_code", None) == 499 and getattr(e, "detail", None) == _CLIENT_DISCONNECT_DETAIL:
|
||||
verbose_proxy_logger.info(
|
||||
"litellm.proxy.proxy_server._handle_llm_api_exception(): client disconnected, upstream LLM request cancelled"
|
||||
"litellm.proxy.proxy_server._handle_llm_api_exception(): client disconnected, "
|
||||
"upstream LLM request cancelled - litellm_call_id=%s",
|
||||
litellm_call_id,
|
||||
extra=MappingProxyType({"litellm_call_id": litellm_call_id}),
|
||||
)
|
||||
return
|
||||
log_fn: Final = (
|
||||
|
|
@ -1462,7 +1466,12 @@ def _log_llm_api_exception(e: Exception) -> None:
|
|||
if is_expected_client_error(e) and not litellm.log_client_error_tracebacks
|
||||
else verbose_proxy_logger.exception
|
||||
)
|
||||
log_fn("litellm.proxy.proxy_server._handle_llm_api_exception(): Exception occured - %s", e)
|
||||
log_fn(
|
||||
"litellm.proxy.proxy_server._handle_llm_api_exception(): Exception occured - litellm_call_id=%s - %s",
|
||||
litellm_call_id,
|
||||
e,
|
||||
extra=MappingProxyType({"litellm_call_id": litellm_call_id}),
|
||||
)
|
||||
|
||||
|
||||
async def _cancel_llm_call_on_client_disconnect(
|
||||
|
|
@ -2338,9 +2347,8 @@ class ProxyBaseLLMRequestProcessing:
|
|||
"""
|
||||
Common request processing logic for both chat completions and responses API endpoints
|
||||
"""
|
||||
requested_model_from_client: Final[str | None] = (
|
||||
self.data.get("model") if isinstance(self.data.get("model"), str) else None
|
||||
)
|
||||
client_model: Final = get_client_requested_model(request) or self.data.get("model")
|
||||
requested_model_from_client: Final[str | None] = client_model if isinstance(client_model, str) else None
|
||||
self._debug_log_request_payload()
|
||||
|
||||
if skip_pre_call_logic:
|
||||
|
|
@ -3421,7 +3429,11 @@ class ProxyBaseLLMRequestProcessing:
|
|||
version: str | None = None,
|
||||
):
|
||||
"""Raises ProxyException (OpenAI API compatible) if an exception is raised"""
|
||||
_log_llm_api_exception(e)
|
||||
logging_obj: Final[LiteLLMLoggingObj | None] = self.data.get("litellm_logging_obj", None)
|
||||
_log_llm_api_exception(
|
||||
e,
|
||||
(logging_obj.litellm_call_id if logging_obj is not None else None) or self.data.get("litellm_call_id"),
|
||||
)
|
||||
# Allow callbacks to transform the error response
|
||||
transformed_exception: Final = await proxy_logging_obj.post_call_failure_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ from fastapi import Request, UploadFile, status
|
|||
from typing_extensions import NotRequired, ReadOnly, Required
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB
|
||||
from litellm.constants import CLIENT_REQUESTED_MODEL_SCOPE_KEY, MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
get_metadata_variable_name_from_kwargs,
|
||||
|
|
@ -235,6 +235,13 @@ def _safe_get_request_parsed_body(request: Request | None) -> dict | None:
|
|||
return None
|
||||
|
||||
|
||||
def get_client_requested_model(request: Request | None) -> str | None:
|
||||
if request is None or not hasattr(request, "scope"):
|
||||
return None
|
||||
model: Final = request.scope.get(CLIENT_REQUESTED_MODEL_SCOPE_KEY)
|
||||
return model if isinstance(model, str) else None
|
||||
|
||||
|
||||
def _safe_get_request_query_params(request: Request | None) -> dict:
|
||||
if request is None:
|
||||
return {}
|
||||
|
|
@ -259,6 +266,24 @@ def _safe_set_request_parsed_body(
|
|||
verbose_proxy_logger.debug("Unexpected error setting request parsed body - %s", e)
|
||||
|
||||
|
||||
def rewrite_request_model(
|
||||
request_data: dict[str, object], # mutable-ok: the request body is rewritten in place for every downstream reader
|
||||
request: Request | None,
|
||||
model: str,
|
||||
) -> None:
|
||||
"""Point the auth-time payload, the parsed-body cache, ``request.json()`` and ``request.body()`` at ``model``.
|
||||
The cache and raw body keep only the keys the client sent, not params auth merged into ``request_data``.
|
||||
"""
|
||||
request_data["model"] = model
|
||||
if request is None:
|
||||
return
|
||||
cached_body: Final = _safe_get_request_parsed_body(request=request)
|
||||
body: Final = {**cached_body, "model": model} if cached_body is not None else request_data
|
||||
_safe_set_request_parsed_body(request=request, parsed_body=body)
|
||||
request._json = body
|
||||
request._body = orjson.dumps(body)
|
||||
|
||||
|
||||
def _safe_get_request_headers(request: Request | None) -> dict:
|
||||
"""
|
||||
[Non-Blocking] Safely get the request headers.
|
||||
|
|
|
|||
|
|
@ -7088,8 +7088,19 @@ class ProxyConfig:
|
|||
|
||||
## PASS-THROUGH ENDPOINTS ##
|
||||
if "pass_through_endpoints" in _general_settings:
|
||||
general_settings["pass_through_endpoints"] = _general_settings["pass_through_endpoints"]
|
||||
await initialize_pass_through_endpoints(pass_through_endpoints=general_settings["pass_through_endpoints"])
|
||||
db_pass_through_endpoints: Final = _general_settings["pass_through_endpoints"]
|
||||
db_pass_through_paths: Final = frozenset(
|
||||
endpoint.get("path") for endpoint in db_pass_through_endpoints if isinstance(endpoint, dict)
|
||||
)
|
||||
general_settings["pass_through_endpoints"] = [
|
||||
*db_pass_through_endpoints,
|
||||
*(
|
||||
endpoint
|
||||
for endpoint in config_passthrough_endpoints or ()
|
||||
if endpoint.get("path") not in db_pass_through_paths
|
||||
),
|
||||
]
|
||||
await initialize_pass_through_endpoints(pass_through_endpoints=db_pass_through_endpoints)
|
||||
|
||||
## UI ACCESS MODE ##
|
||||
if "ui_access_mode" in _general_settings:
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ import subprocess
|
|||
import sys
|
||||
from contextlib import contextmanager
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from functools import partial
|
||||
from pathlib import Path
|
||||
from textwrap import dedent
|
||||
from types import SimpleNamespace
|
||||
|
|
@ -40,6 +41,7 @@ from litellm.proxy.auth.auth_checks import (
|
|||
jwt_key_mapping_cache_key,
|
||||
)
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
from litellm.proxy.common_utils.http_parsing_utils import get_client_requested_model
|
||||
from litellm.proxy.auth.user_api_key_auth import (
|
||||
_check_key_model_budget_with_fallback,
|
||||
_ensure_litellm_received_at_on_request_state,
|
||||
|
|
@ -8298,3 +8300,189 @@ async def test_auth_flow_enters_virtual_key_mapping_when_only_an_issuer_configur
|
|||
assert resolve_mock.await_args.kwargs["jwt_claims"][JWTHandler.LITELLM_JWT_ISSUER_CLAIM] == ISSUER_TWO
|
||||
assert result.api_key == "hashed-mapped-key"
|
||||
assert result.team_id == "svc-team"
|
||||
|
||||
|
||||
def _alias_router() -> litellm.Router:
|
||||
return litellm.Router(
|
||||
model_list=[
|
||||
{"model_name": name, "litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-fake"}}
|
||||
for name in ("claude-haiku", "claude-sonnet")
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def _alias_request(route: str, data: dict, content_type: str = "application/json", path_params: dict | None = None):
|
||||
"""A request as auth sees it: the body already read once and cached alongside its parsed form."""
|
||||
from starlette.requests import Request
|
||||
|
||||
scope = {
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": route,
|
||||
"headers": [(b"content-type", content_type.encode())],
|
||||
"query_string": b"",
|
||||
"path_params": path_params or {},
|
||||
"parsed_body": (tuple(data), data),
|
||||
}
|
||||
request = Request(scope)
|
||||
request._body = json.dumps(data).encode()
|
||||
return request
|
||||
|
||||
|
||||
async def _enforce_alias_access(token: UserAPIKeyAuth, data: dict, route: str, request, router: litellm.Router):
|
||||
from litellm.proxy.auth.user_api_key_auth import _enforce_key_and_fallback_model_access
|
||||
|
||||
await _enforce_key_and_fallback_model_access(
|
||||
valid_token=token,
|
||||
request_data=data,
|
||||
route=route,
|
||||
request=request,
|
||||
llm_model_list=router.model_list,
|
||||
llm_router=router,
|
||||
)
|
||||
|
||||
|
||||
def _alias_token(monkeypatch, level: str, alias: dict, models: list) -> UserAPIKeyAuth:
|
||||
"""A key whose ``router_settings.model_group_alias`` lives on the key itself or on its cached team row."""
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
|
||||
if level == "key":
|
||||
return UserAPIKeyAuth(models=models, router_settings={"model_group_alias": alias})
|
||||
cache = UserApiKeyCache()
|
||||
team = LiteLLM_TeamTableCachedObj(team_id="team-alias", models=models, router_settings={"model_group_alias": alias})
|
||||
cache.set_cache(key="team_id:team-alias", value=team)
|
||||
monkeypatch.setattr(litellm.proxy.proxy_server, "user_api_key_cache", cache)
|
||||
monkeypatch.setattr(litellm.proxy.proxy_server, "prisma_client", MagicMock())
|
||||
return UserAPIKeyAuth(team_id="team-alias", models=models)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("level", ["key", "team"])
|
||||
@pytest.mark.parametrize(
|
||||
"route",
|
||||
["/v1/chat/completions", "/v1/messages", "/v1/embeddings", "/openai/v1/responses", "/cursor/chat/completions"],
|
||||
)
|
||||
async def test_router_settings_model_group_alias_authorizes_target_for_key(monkeypatch, level, route):
|
||||
"""LIT-3054: a key allowed only the alias target must be able to call the alias, and a key not
|
||||
allowed the target must still be denied even when the alias itself is what it requested."""
|
||||
router = _alias_router()
|
||||
monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", router)
|
||||
data = {"model": "AgentX-LLM", "messages": [{"role": "user", "content": "hi"}]}
|
||||
request = _alias_request(route, data)
|
||||
token = _alias_token(monkeypatch, level, {"AgentX-LLM": "claude-haiku"}, ["claude-haiku"])
|
||||
await _enforce_alias_access(token, data, route, request, router)
|
||||
assert data["model"] == "claude-haiku"
|
||||
assert (await request.json())["model"] == "claude-haiku"
|
||||
assert json.loads(await request.body())["model"] == "claude-haiku"
|
||||
assert request.scope["parsed_body"][1]["model"] == "claude-haiku"
|
||||
assert get_client_requested_model(request) == "AgentX-LLM"
|
||||
|
||||
denied = _alias_token(monkeypatch, level, {"AgentX-LLM": "claude-sonnet"}, ["claude-haiku"])
|
||||
denied_data = {"model": "AgentX-LLM"}
|
||||
with pytest.raises(ProxyException) as exc:
|
||||
await _enforce_alias_access(denied, denied_data, route, _alias_request(route, denied_data), router)
|
||||
assert "claude-sonnet" in exc.value.message
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_settings_model_group_alias_leaves_form_bodies_alone(monkeypatch):
|
||||
"""LIT-3054: a multipart body cannot be re-serialized as JSON, so auth must not rewrite it."""
|
||||
router = _alias_router()
|
||||
monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", router)
|
||||
data = {"model": "AgentX-LLM"}
|
||||
route = "/v1/audio/transcriptions"
|
||||
request = _alias_request(route, data, content_type="multipart/form-data; boundary=x")
|
||||
token = _alias_token(monkeypatch, "key", {"AgentX-LLM": "claude-haiku"}, ["claude-haiku", "AgentX-LLM"])
|
||||
await _enforce_alias_access(token, data, route, request, router)
|
||||
assert data["model"] == "AgentX-LLM"
|
||||
assert get_client_requested_model(request) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_settings_model_group_alias_rewrite_keeps_query_params_out_of_body(monkeypatch):
|
||||
"""LIT-3054: auth merges query params into its own copy of the body; the rewrite must not forward them."""
|
||||
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body, populate_request_with_path_params
|
||||
|
||||
router = _alias_router()
|
||||
monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", router)
|
||||
body = {"model": "AgentX-LLM", "messages": [{"role": "user", "content": "hi"}]}
|
||||
request = _alias_request("/v1/chat/completions", body)
|
||||
request.scope["query_string"] = b"api-version=2024-10-21&stream=true"
|
||||
data = populate_request_with_path_params(request_data=await _read_request_body(request), request=request)
|
||||
assert data["api-version"] == "2024-10-21"
|
||||
token = _alias_token(monkeypatch, "key", {"AgentX-LLM": "claude-haiku"}, ["claude-haiku"])
|
||||
await _enforce_alias_access(token, data, "/v1/chat/completions", request, router)
|
||||
downstream = await _read_request_body(request)
|
||||
assert downstream == {**body, "model": "claude-haiku"}
|
||||
assert json.loads(await request.body()) == downstream
|
||||
assert await request.json() == downstream
|
||||
|
||||
|
||||
def _user_defined_pass_through_endpoint():
|
||||
from litellm.types.passthrough_endpoints.pass_through_endpoints import LITELLM_PASS_THROUGH_ENDPOINT_MARKER
|
||||
|
||||
async def endpoint():
|
||||
return None
|
||||
|
||||
setattr(endpoint, LITELLM_PASS_THROUGH_ENDPOINT_MARKER, True)
|
||||
return endpoint
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("user_defined", [False, True])
|
||||
async def test_router_settings_model_group_alias_leaves_pass_through_bodies_alone(monkeypatch, user_defined):
|
||||
"""LIT-3054: pass-through handlers forward the body verbatim to the provider, so auth must not rewrite it.
|
||||
Built-in provider handlers bind ``{endpoint:path}``; user-defined ones carry the pass-through marker."""
|
||||
router = _alias_router()
|
||||
monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", router)
|
||||
data = {"model": "AgentX-LLM", "messages": [{"role": "user", "content": "hi"}]}
|
||||
route = "/custom-upstream/chat" if user_defined else "/anthropic/v1/messages"
|
||||
request = _alias_request(route, data, path_params={} if user_defined else {"endpoint": "v1/messages"})
|
||||
if user_defined:
|
||||
request.scope["endpoint"] = _user_defined_pass_through_endpoint()
|
||||
LiteLLMRoutes.openai_routes.value.append(route)
|
||||
token = _alias_token(monkeypatch, "key", {"AgentX-LLM": "claude-haiku"}, ["claude-haiku", "AgentX-LLM"])
|
||||
try:
|
||||
await _enforce_alias_access(token, data, route, request, router)
|
||||
finally:
|
||||
if user_defined:
|
||||
LiteLLMRoutes.openai_routes.value.remove(route)
|
||||
assert data["model"] == "AgentX-LLM"
|
||||
assert (await request.json())["model"] == "AgentX-LLM"
|
||||
assert get_client_requested_model(request) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("target, expect_denied", [("claude-haiku", False), ("claude-sonnet", True)])
|
||||
async def test_router_settings_model_group_alias_authorizes_target_for_team(monkeypatch, target, expect_denied):
|
||||
"""LIT-3054: the team allowlist check in common_checks must judge the alias target, not the alias."""
|
||||
import litellm.proxy.proxy_server as _proxy_server_mod
|
||||
from litellm.proxy.auth.user_api_key_auth import _authorize_authenticated_request
|
||||
|
||||
router = _alias_router()
|
||||
logging_obj = MagicMock(post_call_failure_hook=AsyncMock(return_value=None))
|
||||
attrs = {**_proxy_attrs_for_centralized_checks(), "llm_router": router, "proxy_logging_obj": logging_obj}
|
||||
for k, v in attrs.items():
|
||||
monkeypatch.setattr(_proxy_server_mod, k, v)
|
||||
token = _alias_token(monkeypatch, "team", {"AgentX-LLM": target}, ["claude-haiku"])
|
||||
token.team_models = ["claude-haiku"]
|
||||
data = {"model": "AgentX-LLM", "messages": [{"role": "user", "content": "hi"}]}
|
||||
route = "/v1/chat/completions"
|
||||
request = _alias_request(route, data)
|
||||
authorize = partial(
|
||||
_authorize_authenticated_request,
|
||||
user_api_key_auth_obj=token,
|
||||
request=request,
|
||||
request_data=data,
|
||||
route=route,
|
||||
api_key="sk-test",
|
||||
)
|
||||
if expect_denied:
|
||||
with pytest.raises(ProxyException) as exc:
|
||||
await authorize()
|
||||
assert exc.value.type == ProxyErrorTypes.team_model_access_denied
|
||||
assert target in exc.value.message
|
||||
return
|
||||
await authorize()
|
||||
assert (await request.json())["model"] == target
|
||||
assert get_client_requested_model(request) == "AgentX-LLM"
|
||||
|
|
|
|||
|
|
@ -13,7 +13,11 @@ from fastapi.responses import JSONResponse, StreamingResponse
|
|||
|
||||
import litellm
|
||||
from litellm._uuid import uuid
|
||||
from litellm.constants import MAX_LITELLM_CALL_ID_LENGTH, RETURN_RAW_MODEL_NAME_METADATA_KEY
|
||||
from litellm.constants import (
|
||||
CLIENT_REQUESTED_MODEL_SCOPE_KEY,
|
||||
MAX_LITELLM_CALL_ID_LENGTH,
|
||||
RETURN_RAW_MODEL_NAME_METADATA_KEY,
|
||||
)
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.integrations.opentelemetry import UserAPIKeyAuth
|
||||
from litellm.proxy.common_request_processing import (
|
||||
|
|
@ -4395,6 +4399,53 @@ class TestDisconnectGatherCleanup:
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("client_model, expected", [("AgentX-LLM", "AgentX-LLM"), (None, "gpt-mini")])
|
||||
async def test_response_model_echoes_the_name_the_client_sent_before_auth_rewrote_it(
|
||||
monkeypatch, client_model, expected
|
||||
):
|
||||
"""LIT-3054: auth resolves router_settings.model_group_alias in the body, so the alias the
|
||||
client sent only survives in the request scope. The response must still echo it."""
|
||||
import litellm.proxy.common_request_processing as cpr
|
||||
|
||||
async def llm():
|
||||
return litellm.ModelResponse(
|
||||
model="gpt-4o-mini", choices=[{"message": {"role": "assistant", "content": "pong"}}]
|
||||
)
|
||||
|
||||
async def fake_route_request(**_kwargs):
|
||||
return llm()
|
||||
|
||||
logging_obj = MagicMock(litellm_call_id="call-id", _defer_async_logging=False)
|
||||
proxy_logging = MagicMock(spec=ProxyLogging)
|
||||
proxy_logging.during_call_hook = AsyncMock(return_value=None)
|
||||
proxy_logging.post_call_success_hook = AsyncMock(side_effect=lambda data, user_api_key_dict, response: response)
|
||||
proxy_logging.post_call_response_headers_hook = AsyncMock(return_value={})
|
||||
proxy_logging._callback_capabilities_cache = {}
|
||||
monkeypatch.setattr(cpr, "route_request", fake_route_request)
|
||||
|
||||
processor = ProxyBaseLLMRequestProcessing(data={"model": "gpt-mini", "messages": []})
|
||||
monkeypatch.setattr(
|
||||
processor, "common_processing_pre_call_logic", AsyncMock(return_value=({"model": "gpt-mini"}, logging_obj))
|
||||
)
|
||||
monkeypatch.setattr(processor, "_has_post_call_guardrails", MagicMock(return_value=False))
|
||||
scope = {"type": "http", "method": "POST", "path": "/v1/chat/completions", "headers": [], "query_string": b""}
|
||||
request = Request({**scope, CLIENT_REQUESTED_MODEL_SCOPE_KEY: client_model} if client_model else scope)
|
||||
|
||||
response = await processor.base_process_llm_request(
|
||||
request=request,
|
||||
fastapi_response=Response(),
|
||||
user_api_key_dict=ProxyUserAPIKeyAuth(),
|
||||
proxy_logging_obj=proxy_logging,
|
||||
general_settings={},
|
||||
proxy_config=MagicMock(spec=ProxyConfig),
|
||||
route_type="acompletion",
|
||||
version=None,
|
||||
)
|
||||
|
||||
assert response.model == expected
|
||||
|
||||
|
||||
class TestStreamingClientDisconnectLogging:
|
||||
@pytest.mark.asyncio
|
||||
async def test_record_streaming_client_disconnect_sets_error_information(self):
|
||||
|
|
@ -8169,7 +8220,7 @@ def test_log_llm_api_exception_traceback_only_for_unexpected_errors(exc, expect_
|
|||
try:
|
||||
raise exc
|
||||
except Exception as raised:
|
||||
_log_llm_api_exception(raised)
|
||||
_log_llm_api_exception(raised, "call-id-for-traceback-test")
|
||||
finally:
|
||||
verbose_proxy_logger.propagate = False
|
||||
|
||||
|
|
@ -8663,3 +8714,84 @@ class TestBackgroundResponseRetrievalGovernance:
|
|||
|
||||
assert "_guardrail_pipelines" not in data["litellm_metadata"]
|
||||
assert "applied_policies" not in data["litellm_metadata"]
|
||||
|
||||
|
||||
class TestErrorLogCarriesCallId:
|
||||
"""Regression for LIT-5856 / #37532: the ERROR line emitted for a failed LLM
|
||||
request must carry the litellm_call_id the client got back in the
|
||||
x-litellm-call-id response header, so a logged exception can be tied to a
|
||||
specific request."""
|
||||
|
||||
async def _invoke(self, data: dict[str, object]) -> None:
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
||||
processor: Final = ProxyBaseLLMRequestProcessing(data=data)
|
||||
proxy_logging_obj: Final = MagicMock()
|
||||
proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
|
||||
verbose_proxy_logger.propagate = True
|
||||
try:
|
||||
with pytest.raises(ProxyException):
|
||||
await processor._handle_llm_api_exception(
|
||||
e=ValueError("upstream blew up"),
|
||||
user_api_key_dict=ProxyUserAPIKeyAuth(api_key="sk-test"),
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
finally:
|
||||
verbose_proxy_logger.propagate = False
|
||||
|
||||
@staticmethod
|
||||
def _error_record(caplog: pytest.LogCaptureFixture):
|
||||
return next(r for r in caplog.records if "_handle_llm_api_exception(): Exception occured" in r.getMessage())
|
||||
|
||||
async def test_call_id_from_logging_obj_is_logged(self, caplog: pytest.LogCaptureFixture) -> None:
|
||||
call_id: Final = str(uuid.uuid4())
|
||||
logging_obj: Final = MagicMock()
|
||||
logging_obj.litellm_call_id = call_id
|
||||
with caplog.at_level("ERROR", logger="LiteLLM Proxy"):
|
||||
await self._invoke({"litellm_logging_obj": logging_obj, "litellm_call_id": "stale-id"})
|
||||
|
||||
record: Final = self._error_record(caplog)
|
||||
assert record.litellm_call_id == call_id
|
||||
assert call_id in record.getMessage()
|
||||
|
||||
async def test_call_id_falls_back_to_request_data(self, caplog: pytest.LogCaptureFixture) -> None:
|
||||
call_id: Final = str(uuid.uuid4())
|
||||
with caplog.at_level("ERROR", logger="LiteLLM Proxy"):
|
||||
await self._invoke({"litellm_call_id": call_id})
|
||||
|
||||
record: Final = self._error_record(caplog)
|
||||
assert record.litellm_call_id == call_id
|
||||
assert call_id in record.getMessage()
|
||||
|
||||
async def test_call_id_falls_back_when_logging_obj_has_none(self, caplog: pytest.LogCaptureFixture) -> None:
|
||||
call_id: Final = str(uuid.uuid4())
|
||||
logging_obj: Final = MagicMock()
|
||||
logging_obj.litellm_call_id = None
|
||||
with caplog.at_level("ERROR", logger="LiteLLM Proxy"):
|
||||
await self._invoke({"litellm_logging_obj": logging_obj, "litellm_call_id": call_id})
|
||||
|
||||
record: Final = self._error_record(caplog)
|
||||
assert record.litellm_call_id == call_id
|
||||
assert call_id in record.getMessage()
|
||||
|
||||
def test_client_disconnect_log_carries_call_id(self, caplog: pytest.LogCaptureFixture) -> None:
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy.common_request_processing import (
|
||||
_CLIENT_DISCONNECT_DETAIL,
|
||||
_log_llm_api_exception,
|
||||
)
|
||||
|
||||
call_id: Final = str(uuid.uuid4())
|
||||
verbose_proxy_logger.propagate = True
|
||||
try:
|
||||
with caplog.at_level("INFO", logger="LiteLLM Proxy"):
|
||||
_log_llm_api_exception(
|
||||
HTTPException(status_code=499, detail=_CLIENT_DISCONNECT_DETAIL),
|
||||
call_id,
|
||||
)
|
||||
finally:
|
||||
verbose_proxy_logger.propagate = False
|
||||
|
||||
record: Final = caplog.records[-1]
|
||||
assert record.litellm_call_id == call_id
|
||||
assert call_id in record.getMessage()
|
||||
|
|
|
|||
|
|
@ -7408,6 +7408,88 @@ async def test_update_general_settings_apply_user_budget_to_team_keys_yaml_wins(
|
|||
assert ps.general_settings["apply_user_budget_to_team_keys"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_general_settings_keeps_yaml_pass_through_endpoints_next_to_db_ones():
|
||||
"""user_api_key_auth honours ``auth: false`` only for entries it finds in
|
||||
general_settings["pass_through_endpoints"]. The DB overlay used to replace that
|
||||
list wholesale, so once one endpoint existed in the DB the YAML-declared
|
||||
auth-disabled route started answering 401 while staying registered."""
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.proxy_server import ProxyConfig
|
||||
|
||||
yaml_endpoint: Final = {"path": "/v1/cuopt/request", "target": "https://example.com/post", "auth": False}
|
||||
db_endpoint: Final = {"id": "db-1", "path": "/v1/db-echo", "target": "https://example.com/post", "auth": True}
|
||||
|
||||
def request_without_key(path: str) -> MagicMock:
|
||||
request: Final = MagicMock()
|
||||
request.url.path = path
|
||||
request.headers = {}
|
||||
request.query_params = {}
|
||||
return request
|
||||
|
||||
settings: Final = patch("litellm.proxy.proxy_server.general_settings", {"pass_through_endpoints": [yaml_endpoint]}) # test-quality-ok: the method reads this module global; no injection seam
|
||||
yaml_endpoints: Final = patch("litellm.proxy.proxy_server.config_passthrough_endpoints", [yaml_endpoint]) # test-quality-ok: module global holding the YAML endpoints the fix merges in
|
||||
initialize: Final = patch("litellm.proxy.proxy_server.initialize_pass_through_endpoints", AsyncMock()) # test-quality-ok: route registration needs the FastAPI app; auth is the observable here
|
||||
master_key: Final = patch("litellm.proxy.proxy_server.master_key", "sk-master") # test-quality-ok: a set master key is what makes a missing Authorization header a 401
|
||||
with settings, yaml_endpoints, initialize, master_key:
|
||||
await ProxyConfig()._update_general_settings(db_general_settings={"pass_through_endpoints": [db_endpoint]})
|
||||
|
||||
anonymous: Final = await user_api_key_auth(request=request_without_key("/v1/cuopt/request"), api_key=None)
|
||||
assert anonymous.api_key is None
|
||||
|
||||
with pytest.raises(ProxyException) as still_protected:
|
||||
await user_api_key_auth(request=request_without_key("/v1/db-echo"), api_key=None)
|
||||
assert still_protected.value.code == "401"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("db_methods", "yaml_methods"),
|
||||
[(None, None), (["POST"], ["GET"])],
|
||||
ids=["all-methods", "disjoint-methods"],
|
||||
)
|
||||
async def test_update_general_settings_db_pass_through_endpoint_overrides_yaml_entry_on_the_same_path(
|
||||
db_methods: list[str] | None, yaml_methods: list[str] | None
|
||||
):
|
||||
"""The auth check matches pass-through entries by path only and lets any
|
||||
matching ``auth: false`` entry through, so a DB ``auth: true`` entry can only
|
||||
lock down a YAML-declared path if the YAML entry is dropped from the merged
|
||||
list, whatever ``methods`` either entry declares."""
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.proxy_server import ProxyConfig
|
||||
|
||||
yaml_endpoint: Final = {
|
||||
"path": "/v1/cuopt/request",
|
||||
"target": "https://example.com/post",
|
||||
"auth": False,
|
||||
"methods": yaml_methods,
|
||||
}
|
||||
db_endpoint: Final = {
|
||||
"id": "db-1",
|
||||
"path": "/v1/cuopt/request",
|
||||
"target": "https://example.com/post",
|
||||
"auth": True,
|
||||
"methods": db_methods,
|
||||
}
|
||||
|
||||
request: Final = MagicMock()
|
||||
request.url.path = "/v1/cuopt/request"
|
||||
request.method = "POST"
|
||||
request.headers = {}
|
||||
request.query_params = {}
|
||||
|
||||
settings: Final = patch("litellm.proxy.proxy_server.general_settings", {"pass_through_endpoints": [yaml_endpoint]}) # test-quality-ok: the method reads this module global; no injection seam
|
||||
yaml_endpoints: Final = patch("litellm.proxy.proxy_server.config_passthrough_endpoints", [yaml_endpoint]) # test-quality-ok: module global holding the YAML endpoints the fix merges in
|
||||
initialize: Final = patch("litellm.proxy.proxy_server.initialize_pass_through_endpoints", AsyncMock()) # test-quality-ok: route registration needs the FastAPI app; auth is the observable here
|
||||
master_key: Final = patch("litellm.proxy.proxy_server.master_key", "sk-master") # test-quality-ok: a set master key is what makes a missing Authorization header a 401
|
||||
with settings, yaml_endpoints, initialize, master_key:
|
||||
await ProxyConfig()._update_general_settings(db_general_settings={"pass_through_endpoints": [db_endpoint]})
|
||||
|
||||
with pytest.raises(ProxyException) as locked_down:
|
||||
await user_api_key_auth(request=request, api_key=None)
|
||||
assert locked_down.value.code == "401"
|
||||
|
||||
|
||||
def _fill_user_api_key_cache(cache: DualCache, count: int) -> None:
|
||||
for index in range(count):
|
||||
cache.set_cache(key=f"key-{index}", value={"token": f"key-{index}"}, local_only=True)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue