mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-08 22:21:35 +00:00
fix(proxy): route invalid virtual key logs to stdout
This commit is contained in:
parent
4d9025c2bf
commit
1a1d6ce804
7 changed files with 481 additions and 52 deletions
|
|
@ -264,13 +264,16 @@ def _plain_log_format(stdout: TextIO | None, stderr: TextIO | None) -> str:
|
|||
|
||||
|
||||
class LevelRoutingStreamHandler(logging.StreamHandler):
|
||||
"""Writes records below WARNING to stdout and WARNING and above to stderr.
|
||||
"""Writes records below WARNING and invalid-key warnings to stdout.
|
||||
|
||||
Collectors that derive severity from the stream report every stderr line as an error.
|
||||
"""
|
||||
|
||||
def emit(self, record: logging.LogRecord) -> None:
|
||||
preferred: Final = sys.stdout if record.levelno < logging.WARNING else sys.stderr
|
||||
is_stdout_record: Final = record.levelno < logging.WARNING or (
|
||||
record.levelno == logging.WARNING and record.name == verbose_proxy_stdout_logger.name
|
||||
)
|
||||
preferred: Final = sys.stdout if is_stdout_record else sys.stderr
|
||||
if preferred is None or getattr(preferred, "closed", False):
|
||||
self.stream = sys.stderr # rebind-ok: fall back to the pre-fix stream rather than raising per record
|
||||
else:
|
||||
|
|
@ -507,19 +510,30 @@ else:
|
|||
|
||||
handler.setFormatter(formatter)
|
||||
|
||||
verbose_proxy_logger = logging.getLogger("LiteLLM Proxy")
|
||||
verbose_router_logger = logging.getLogger("LiteLLM Router")
|
||||
verbose_logger = logging.getLogger("LiteLLM")
|
||||
verbose_proxy_logger: Final = logging.getLogger("LiteLLM Proxy")
|
||||
verbose_proxy_stdout_logger: Final = verbose_proxy_logger.getChild("stdout")
|
||||
verbose_router_logger: Final = logging.getLogger("LiteLLM Router")
|
||||
verbose_logger: Final = logging.getLogger("LiteLLM")
|
||||
|
||||
verbose_proxy_stdout_handler: Final = LevelRoutingStreamHandler()
|
||||
verbose_proxy_stdout_handler.setLevel(logging.WARNING)
|
||||
verbose_proxy_stdout_handler.setFormatter(handler.formatter)
|
||||
verbose_proxy_stdout_handler.addFilter(_secret_filter)
|
||||
verbose_proxy_stdout_handler.addFilter(_correlation_filter)
|
||||
|
||||
# Add the handler to the loggers
|
||||
verbose_router_logger.addHandler(handler)
|
||||
verbose_proxy_logger.addHandler(handler)
|
||||
verbose_proxy_stdout_logger.setLevel(logging.WARNING)
|
||||
verbose_proxy_stdout_logger.addHandler(verbose_proxy_stdout_handler)
|
||||
verbose_proxy_stdout_logger.propagate = False
|
||||
verbose_logger.addHandler(handler)
|
||||
|
||||
# Filters attached to the logger, not the handler, survive callers swapping in their own
|
||||
# handlers (JSON mode, uvicorn log config, a host app's root handler).
|
||||
verbose_router_logger.addFilter(_stdout_truncation_filter)
|
||||
verbose_proxy_logger.addFilter(_stdout_truncation_filter)
|
||||
verbose_proxy_stdout_logger.addFilter(_stdout_truncation_filter)
|
||||
verbose_logger.addFilter(_stdout_truncation_filter)
|
||||
|
||||
|
||||
|
|
@ -578,6 +592,7 @@ ALL_LOGGERS: Final = [
|
|||
verbose_logger,
|
||||
verbose_router_logger,
|
||||
verbose_proxy_logger,
|
||||
verbose_proxy_stdout_logger,
|
||||
]
|
||||
|
||||
|
||||
|
|
@ -685,6 +700,7 @@ def _turn_on_json():
|
|||
handler: Final = LevelRoutingStreamHandler()
|
||||
handler.setFormatter(JsonFormatter())
|
||||
_initialize_loggers_with_handler(handler)
|
||||
verbose_proxy_stdout_logger.setLevel(logging.WARNING)
|
||||
# Set up exception handlers
|
||||
_setup_json_exception_handlers(JsonFormatter())
|
||||
|
||||
|
|
@ -700,12 +716,14 @@ def _disable_debugging():
|
|||
verbose_logger.disabled = True
|
||||
verbose_router_logger.disabled = True
|
||||
verbose_proxy_logger.disabled = True
|
||||
verbose_proxy_stdout_logger.disabled = True
|
||||
|
||||
|
||||
def _enable_debugging():
|
||||
verbose_logger.disabled = False
|
||||
verbose_router_logger.disabled = False
|
||||
verbose_proxy_logger.disabled = False
|
||||
verbose_proxy_stdout_logger.disabled = False
|
||||
|
||||
|
||||
def print_verbose(print_statement):
|
||||
|
|
|
|||
|
|
@ -2,13 +2,14 @@
|
|||
Handles Authentication Errors
|
||||
"""
|
||||
|
||||
import logging
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from fastapi import HTTPException, Request, status
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._logging import verbose_proxy_logger, verbose_proxy_stdout_logger
|
||||
from litellm.constants import EMPTY_MAPPING
|
||||
from litellm.integrations.otel.runtime import seed_request_identity
|
||||
from litellm.litellm_core_utils.core_helpers import is_expected_client_error
|
||||
|
|
@ -18,7 +19,11 @@ from litellm.proxy._types import (
|
|||
ProxyException,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.auth.auth_utils import _get_request_ip_address
|
||||
from litellm.proxy.auth.auth_utils import (
|
||||
_get_request_ip_address,
|
||||
is_invalid_virtual_key_error,
|
||||
mark_invalid_virtual_key_error,
|
||||
)
|
||||
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
|
||||
from litellm.types.services import ServiceTypes
|
||||
|
||||
|
|
@ -110,17 +115,8 @@ class UserAPIKeyAuthExceptionHandler:
|
|||
request=request,
|
||||
use_x_forwarded_for=general_settings.get("use_x_forwarded_for") is True,
|
||||
)
|
||||
log_fn: Final = (
|
||||
verbose_proxy_logger.error
|
||||
if is_expected_client_error(e) and not litellm.log_client_error_tracebacks
|
||||
else verbose_proxy_logger.exception
|
||||
)
|
||||
log_fn(
|
||||
"litellm.proxy.proxy_server.user_api_key_auth(): Exception occured - %s\nRequester IP Address:%s",
|
||||
e,
|
||||
requester_ip,
|
||||
extra={"requester_ip": requester_ip},
|
||||
)
|
||||
original_exception: Final = e
|
||||
is_invalid_virtual_key: Final = is_invalid_virtual_key_error(e)
|
||||
|
||||
# Log this exception to OTEL, Datadog etc. Reuse the identity resolved
|
||||
# before the failure (team alias/id, metadata, user) so the failed span
|
||||
|
|
@ -158,7 +154,7 @@ class UserAPIKeyAuthExceptionHandler:
|
|||
# Allow callbacks to transform the error response
|
||||
transformed_exception: Final = await proxy_logging_obj.post_call_failure_hook(
|
||||
request_data=_with_requester_ip_address(request_data, requester_ip),
|
||||
original_exception=e,
|
||||
original_exception=original_exception,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
error_type=ProxyErrorTypes.auth_error,
|
||||
route=route,
|
||||
|
|
@ -167,24 +163,25 @@ class UserAPIKeyAuthExceptionHandler:
|
|||
if transformed_exception is not None:
|
||||
e = transformed_exception
|
||||
|
||||
proxy_exception: Final
|
||||
if isinstance(e, litellm.BudgetExceededError):
|
||||
raise ProxyException(
|
||||
proxy_exception = ProxyException(
|
||||
message=e.message,
|
||||
type=ProxyErrorTypes.budget_exceeded,
|
||||
param=None,
|
||||
code=getattr(e, "status_code", status.HTTP_429_TOO_MANY_REQUESTS),
|
||||
)
|
||||
if isinstance(e, HTTPException):
|
||||
raise ProxyException(
|
||||
elif isinstance(e, HTTPException):
|
||||
proxy_exception = ProxyException(
|
||||
message=getattr(e, "detail", f"Authentication Error({e})"),
|
||||
type=ProxyErrorTypes.auth_error,
|
||||
param=getattr(e, "param", "None"),
|
||||
code=getattr(e, "status_code", status.HTTP_401_UNAUTHORIZED),
|
||||
)
|
||||
elif isinstance(e, ProxyException):
|
||||
raise e
|
||||
if PrismaDBExceptionHandler.is_database_service_unavailable_error(e):
|
||||
raise ProxyException(
|
||||
proxy_exception = e
|
||||
elif PrismaDBExceptionHandler.is_database_service_unavailable_error(e):
|
||||
proxy_exception = ProxyException(
|
||||
message=(
|
||||
"Service Unavailable, the authentication database is "
|
||||
"temporarily unreachable. Please retry shortly."
|
||||
|
|
@ -193,9 +190,24 @@ class UserAPIKeyAuthExceptionHandler:
|
|||
param="None",
|
||||
code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
)
|
||||
raise ProxyException(
|
||||
message="Authentication Error, " + str(e),
|
||||
type=ProxyErrorTypes.auth_error,
|
||||
param=getattr(e, "param", "None"),
|
||||
code=status.HTTP_401_UNAUTHORIZED,
|
||||
else:
|
||||
proxy_exception = ProxyException(
|
||||
message="Authentication Error, " + str(e),
|
||||
type=ProxyErrorTypes.auth_error,
|
||||
param=getattr(e, "param", "None"),
|
||||
code=status.HTTP_401_UNAUTHORIZED,
|
||||
)
|
||||
final_exception: Final = mark_invalid_virtual_key_error(proxy_exception, is_invalid_virtual_key)
|
||||
is_quiet_log: Final = (
|
||||
is_invalid_virtual_key_error(final_exception) and not litellm.log_client_error_tracebacks
|
||||
)
|
||||
logger: Final = verbose_proxy_stdout_logger if is_quiet_log else verbose_proxy_logger
|
||||
logger.log(
|
||||
logging.WARNING if is_quiet_log else logging.ERROR,
|
||||
"litellm.proxy.proxy_server.user_api_key_auth(): Exception occured - %s\nRequester IP Address:%s",
|
||||
final_exception,
|
||||
requester_ip,
|
||||
exc_info=not (is_expected_client_error(original_exception) and not litellm.log_client_error_tracebacks),
|
||||
extra={"requester_ip": requester_ip},
|
||||
)
|
||||
raise final_exception
|
||||
|
|
|
|||
|
|
@ -33,6 +33,9 @@ from litellm.types.passthrough_endpoints.pass_through_endpoints import (
|
|||
from litellm.types.router import CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS
|
||||
from litellm.types.utils import CustomPricingLiteLLMParams
|
||||
|
||||
INVALID_VIRTUAL_KEY_ERROR_MESSAGE: Final = "LiteLLM Virtual Key expected"
|
||||
_INVALID_VIRTUAL_KEY_ERROR_MARKER: Final = "_litellm_invalid_virtual_key_error"
|
||||
|
||||
|
||||
def _get_request_ip_address(request: Request, use_x_forwarded_for: bool | None = False) -> str | None:
|
||||
client_ip = None
|
||||
|
|
@ -46,6 +49,40 @@ def _get_request_ip_address(request: Request, use_x_forwarded_for: bool | None =
|
|||
return client_ip
|
||||
|
||||
|
||||
def is_invalid_virtual_key_error(exception: BaseException | None) -> bool:
|
||||
"""True when an authentication error rejects a malformed virtual key."""
|
||||
if not isinstance(exception, (HTTPException, ProxyException)):
|
||||
return False
|
||||
|
||||
code: Final[object] = getattr(exception, "code", None)
|
||||
status_code: Final[object] = code if code is not None else getattr(exception, "status_code", None)
|
||||
if str(status_code) != str(status.HTTP_401_UNAUTHORIZED):
|
||||
return False
|
||||
|
||||
if getattr(exception, _INVALID_VIRTUAL_KEY_ERROR_MARKER, False) is True:
|
||||
return True
|
||||
|
||||
message: Final = getattr(exception, "detail", None) or getattr(exception, "message", "")
|
||||
return INVALID_VIRTUAL_KEY_ERROR_MESSAGE in str(message)
|
||||
|
||||
|
||||
def mark_invalid_virtual_key_error(exception: ProxyException, is_invalid_virtual_key: bool) -> ProxyException:
|
||||
"""Return an independently marked malformed-key exception after callback transformations."""
|
||||
if not is_invalid_virtual_key or str(exception.code) != str(status.HTTP_401_UNAUTHORIZED):
|
||||
return exception
|
||||
marked_exception: Final = ProxyException(
|
||||
message=exception.message,
|
||||
type=exception.type,
|
||||
param=exception.param,
|
||||
code=exception.code,
|
||||
headers=exception.headers.copy(),
|
||||
openai_code=exception.openai_code,
|
||||
provider_specific_fields=exception.provider_specific_fields,
|
||||
)
|
||||
setattr(marked_exception, _INVALID_VIRTUAL_KEY_ERROR_MARKER, True)
|
||||
return marked_exception
|
||||
|
||||
|
||||
def _check_valid_ip(
|
||||
allowed_ips: list[str] | None,
|
||||
request: Request,
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ import fastapi
|
|||
import orjson
|
||||
from fastapi import HTTPException, Request, WebSocket, status
|
||||
from fastapi.security.api_key import APIKeyHeader
|
||||
from starlette.exceptions import WebSocketException
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger, verbose_proxy_logger
|
||||
|
|
@ -60,11 +61,13 @@ from litellm.proxy.auth.auth_checks import (
|
|||
from litellm.proxy.auth.auth_exception_handler import UserAPIKeyAuthExceptionHandler
|
||||
from litellm.proxy.auth.auth_method import AuthMethod
|
||||
from litellm.proxy.auth.auth_utils import (
|
||||
INVALID_VIRTUAL_KEY_ERROR_MESSAGE,
|
||||
abbreviate_api_key,
|
||||
get_end_user_id_from_request_body,
|
||||
get_model_from_request,
|
||||
get_request_route,
|
||||
get_request_route_template,
|
||||
is_invalid_virtual_key_error,
|
||||
iter_request_fallback_targets,
|
||||
normalize_request_route,
|
||||
pre_db_read_auth_checks,
|
||||
|
|
@ -539,6 +542,8 @@ async def user_api_key_auth_websocket(websocket: WebSocket):
|
|||
try:
|
||||
return await user_api_key_auth(request=request, api_key=f"Bearer {api_key}")
|
||||
except Exception as e:
|
||||
if is_invalid_virtual_key_error(e):
|
||||
raise WebSocketException(code=status.WS_1008_POLICY_VIOLATION)
|
||||
verbose_proxy_logger.exception(e)
|
||||
await websocket.close(code=status.WS_1008_POLICY_VIOLATION)
|
||||
raise HTTPException(status_code=403, detail=str(e))
|
||||
|
|
@ -1870,7 +1875,7 @@ async def _user_api_key_auth_builder(
|
|||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=(
|
||||
f"LiteLLM Virtual Key expected. Received={_masked_key}, "
|
||||
f"{INVALID_VIRTUAL_KEY_ERROR_MESSAGE}. Received={_masked_key}, "
|
||||
f"expected to start with 'sk-'.{_hint}"
|
||||
),
|
||||
) # prevent token hashes from being used
|
||||
|
|
|
|||
|
|
@ -874,6 +874,44 @@ async def test_user_api_key_auth_websocket():
|
|||
assert mock_user_api_key_auth.call_args.kwargs["api_key"] == "Bearer some_api_key"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_user_api_key_auth_websocket_skips_duplicate_invalid_key_log():
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.auth.auth_utils import mark_invalid_virtual_key_error
|
||||
from litellm.proxy.auth.user_api_key_auth import WebSocketException, user_api_key_auth_websocket
|
||||
|
||||
mock_websocket = MagicMock(spec=WebSocket)
|
||||
mock_websocket.query_params = {"model": "some_model"}
|
||||
mock_websocket.headers = {"authorization": "Bearer undefined"}
|
||||
mock_websocket.scope = {"headers": [(b"authorization", b"Bearer undefined")]}
|
||||
mock_websocket.url = URL(url="/v1/responses")
|
||||
transformed_invalid_virtual_key = mark_invalid_virtual_key_error(
|
||||
ProxyException(
|
||||
message="Please authenticate again",
|
||||
type="auth_error",
|
||||
param="None",
|
||||
code=status.HTTP_401_UNAUTHORIZED,
|
||||
),
|
||||
True,
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.user_api_key_auth",
|
||||
side_effect=transformed_invalid_virtual_key,
|
||||
autospec=True,
|
||||
),
|
||||
patch("litellm.proxy.auth.user_api_key_auth.verbose_proxy_logger.exception") as exception_log,
|
||||
pytest.raises(WebSocketException) as exc_info,
|
||||
):
|
||||
await user_api_key_auth_websocket(mock_websocket)
|
||||
|
||||
assert exc_info.value.code == status.WS_1008_POLICY_VIOLATION
|
||||
assert exc_info.value.reason == ""
|
||||
exception_log.assert_not_called()
|
||||
mock_websocket.close.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_user_api_key_auth_websocket_carries_asgi_path():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
|
|
@ -24,8 +25,7 @@ from prisma.errors import (
|
|||
UniqueViolationError,
|
||||
)
|
||||
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._logging import verbose_proxy_logger, verbose_proxy_stdout_logger
|
||||
from litellm.exceptions import BudgetExceededError
|
||||
from litellm.proxy._types import ProxyErrorTypes, ProxyException, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.auth_exception_handler import UserAPIKeyAuthExceptionHandler
|
||||
|
|
@ -335,9 +335,7 @@ async def test_handle_authentication_error_budget_exceeded():
|
|||
# Test with budget exceeded error
|
||||
from litellm.exceptions import BudgetExceededError
|
||||
|
||||
budget_error = BudgetExceededError(
|
||||
message="Budget exceeded", current_cost=100, max_budget=100
|
||||
)
|
||||
budget_error = BudgetExceededError(message="Budget exceeded", current_cost=100, max_budget=100)
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await handler._handle_authentication_error(
|
||||
|
|
@ -705,19 +703,79 @@ async def test_auth_failure_ip_stamp_does_not_mutate_callers_request_data():
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"auth_error,expect_traceback",
|
||||
"auth_error,expected_level,expect_traceback,log_client_error_tracebacks",
|
||||
[
|
||||
pytest.param(
|
||||
ProxyException(
|
||||
message="Authentication Error", type=ProxyErrorTypes.auth_error, param=None, code=401
|
||||
HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="LiteLLM Virtual Key expected. Received=unde****ined, expected to start with 'sk-'.",
|
||||
),
|
||||
logging.WARNING,
|
||||
False,
|
||||
id="expected_401_no_traceback",
|
||||
False,
|
||||
id="invalid_virtual_key_logs_at_warning_without_traceback",
|
||||
),
|
||||
pytest.param(
|
||||
HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="LiteLLM Virtual Key expected. Received=unde****ined, expected to start with 'sk-'.",
|
||||
),
|
||||
logging.ERROR,
|
||||
True,
|
||||
True,
|
||||
id="invalid_virtual_key_keeps_traceback_opt_in",
|
||||
),
|
||||
pytest.param(
|
||||
HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Authentication Error",
|
||||
),
|
||||
logging.ERROR,
|
||||
False,
|
||||
False,
|
||||
id="other_http_401_keeps_error_level_without_traceback",
|
||||
),
|
||||
pytest.param(
|
||||
ProxyException(
|
||||
message="LiteLLM Virtual Key expected",
|
||||
type=ProxyErrorTypes.auth_error,
|
||||
param=None,
|
||||
code=401,
|
||||
),
|
||||
logging.WARNING,
|
||||
False,
|
||||
False,
|
||||
id="custom_auth_invalid_virtual_key_logs_at_warning_without_traceback",
|
||||
),
|
||||
pytest.param(
|
||||
ProxyException(message="Authentication Error", type=ProxyErrorTypes.auth_error, param=None, code=401),
|
||||
logging.ERROR,
|
||||
False,
|
||||
False,
|
||||
id="other_expected_401_keeps_error_level_without_traceback",
|
||||
),
|
||||
pytest.param(
|
||||
ValueError("unexpected internal error"),
|
||||
logging.ERROR,
|
||||
True,
|
||||
False,
|
||||
id="unexpected_error_keeps_traceback",
|
||||
),
|
||||
pytest.param(
|
||||
HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="Authentication database unavailable",
|
||||
),
|
||||
logging.ERROR,
|
||||
True,
|
||||
False,
|
||||
id="server_error_keeps_traceback",
|
||||
),
|
||||
pytest.param(ValueError("unexpected internal error"), True, id="unexpected_error_keeps_traceback"),
|
||||
],
|
||||
)
|
||||
async def test_handle_authentication_error_traceback_only_for_unexpected_errors(auth_error, expect_traceback, caplog):
|
||||
async def test_handle_authentication_error_traceback_only_for_unexpected_errors(
|
||||
auth_error, expected_level, expect_traceback, log_client_error_tracebacks, caplog
|
||||
):
|
||||
"""Regression for LIT-6043: expected 4xx auth rejections must not format a
|
||||
traceback via logger.exception; unexpected errors must keep it."""
|
||||
handler = UserAPIKeyAuthExceptionHandler()
|
||||
|
|
@ -731,17 +789,22 @@ async def test_handle_authentication_error_traceback_only_for_unexpected_errors(
|
|||
patch( # test-quality-ok: handler reads proxy_server globals at call time
|
||||
"litellm.proxy.auth.auth_exception_handler.seed_request_identity",
|
||||
),
|
||||
patch( # test-quality-ok: handler reads this global flag at call time
|
||||
"litellm.proxy.auth.auth_exception_handler.litellm.log_client_error_tracebacks",
|
||||
log_client_error_tracebacks,
|
||||
),
|
||||
patch( # test-quality-ok: handler reads proxy_server globals at call time
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{"allow_requests_on_db_unavailable": False},
|
||||
),
|
||||
):
|
||||
verbose_proxy_logger.propagate = True
|
||||
verbose_proxy_stdout_logger.propagate = True
|
||||
try:
|
||||
try:
|
||||
raise auth_error
|
||||
except (ProxyException, ValueError) as caught:
|
||||
with caplog.at_level("ERROR", logger="LiteLLM Proxy"), pytest.raises(ProxyException):
|
||||
except (HTTPException, ProxyException, ValueError) as caught:
|
||||
with caplog.at_level(logging.DEBUG), pytest.raises(ProxyException):
|
||||
await handler._handle_authentication_error(
|
||||
caught,
|
||||
MagicMock(),
|
||||
|
|
@ -752,7 +815,158 @@ async def test_handle_authentication_error_traceback_only_for_unexpected_errors(
|
|||
)
|
||||
finally:
|
||||
verbose_proxy_logger.propagate = False
|
||||
verbose_proxy_stdout_logger.propagate = False
|
||||
|
||||
records = [r for r in caplog.records if "user_api_key_auth(): Exception occured" in r.getMessage()]
|
||||
assert len(records) == 1
|
||||
assert (records[0].exc_info is not None) is expect_traceback
|
||||
assert records[0].levelno == expected_level
|
||||
assert bool(records[0].exc_info) is expect_traceback
|
||||
assert records[0].name == (
|
||||
verbose_proxy_stdout_logger.name
|
||||
if expected_level == logging.WARNING and not log_client_error_tracebacks
|
||||
else verbose_proxy_logger.name
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"transformed_exception",
|
||||
[
|
||||
HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Please authenticate again"),
|
||||
ProxyException(
|
||||
message="Please authenticate again",
|
||||
type=ProxyErrorTypes.auth_error,
|
||||
param=None,
|
||||
code=status.HTTP_401_UNAUTHORIZED,
|
||||
),
|
||||
],
|
||||
)
|
||||
async def test_handle_authentication_error_preserves_invalid_virtual_key_marker_after_callback_transform(
|
||||
caplog,
|
||||
transformed_exception,
|
||||
):
|
||||
handler = UserAPIKeyAuthExceptionHandler()
|
||||
original_exception = HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="LiteLLM Virtual Key expected. Received=unde****ined, expected to start with 'sk-'.",
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.proxy_logging_obj.post_call_failure_hook",
|
||||
new_callable=AsyncMock,
|
||||
return_value=transformed_exception,
|
||||
),
|
||||
patch("litellm.proxy.auth.auth_exception_handler.seed_request_identity"),
|
||||
patch("litellm.proxy.proxy_server.general_settings", {"allow_requests_on_db_unavailable": False}),
|
||||
):
|
||||
verbose_proxy_stdout_logger.propagate = True
|
||||
try:
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await handler._handle_authentication_error(
|
||||
original_exception,
|
||||
MagicMock(),
|
||||
{},
|
||||
"/v1/chat/completions",
|
||||
None,
|
||||
"undefined",
|
||||
)
|
||||
finally:
|
||||
verbose_proxy_stdout_logger.propagate = False
|
||||
|
||||
records = [r for r in caplog.records if "user_api_key_auth(): Exception occured" in r.getMessage()]
|
||||
assert len(records) == 1
|
||||
assert records[0].name == verbose_proxy_stdout_logger.name
|
||||
assert records[0].levelno == logging.WARNING
|
||||
assert exc_info.value.message == "Please authenticate again"
|
||||
assert getattr(exc_info.value, "_litellm_invalid_virtual_key_error") is True
|
||||
if isinstance(transformed_exception, ProxyException):
|
||||
assert not hasattr(transformed_exception, "_litellm_invalid_virtual_key_error")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handle_authentication_error_keeps_unexpected_source_traceback_after_callback_4xx(
|
||||
caplog,
|
||||
):
|
||||
handler = UserAPIKeyAuthExceptionHandler()
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.proxy_logging_obj.post_call_failure_hook",
|
||||
new_callable=AsyncMock,
|
||||
return_value=HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Please authenticate again",
|
||||
),
|
||||
),
|
||||
patch("litellm.proxy.auth.auth_exception_handler.seed_request_identity"),
|
||||
patch("litellm.proxy.proxy_server.general_settings", {"allow_requests_on_db_unavailable": False}),
|
||||
):
|
||||
verbose_proxy_logger.propagate = True
|
||||
try:
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await handler._handle_authentication_error(
|
||||
ValueError("unexpected internal error"),
|
||||
MagicMock(),
|
||||
{},
|
||||
"/v1/chat/completions",
|
||||
None,
|
||||
"sk-bad-key",
|
||||
)
|
||||
finally:
|
||||
verbose_proxy_logger.propagate = False
|
||||
|
||||
records = [r for r in caplog.records if "user_api_key_auth(): Exception occured" in r.getMessage()]
|
||||
assert len(records) == 1
|
||||
assert records[0].name == verbose_proxy_logger.name
|
||||
assert records[0].levelno == logging.ERROR
|
||||
assert records[0].exc_info is not None
|
||||
assert "Please authenticate again" in records[0].getMessage()
|
||||
assert exc_info.value.code == str(status.HTTP_401_UNAUTHORIZED)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handle_authentication_error_does_not_preserve_invalid_virtual_key_marker_for_callback_503(
|
||||
caplog,
|
||||
):
|
||||
handler = UserAPIKeyAuthExceptionHandler()
|
||||
original_exception = HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="LiteLLM Virtual Key expected. Received=unde****ined, expected to start with 'sk-'.",
|
||||
)
|
||||
transformed_exception = HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="Authentication service temporarily unavailable",
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.proxy_logging_obj.post_call_failure_hook",
|
||||
new_callable=AsyncMock,
|
||||
return_value=transformed_exception,
|
||||
),
|
||||
patch("litellm.proxy.auth.auth_exception_handler.seed_request_identity"),
|
||||
patch("litellm.proxy.proxy_server.general_settings", {"allow_requests_on_db_unavailable": False}),
|
||||
):
|
||||
verbose_proxy_logger.propagate = True
|
||||
try:
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await handler._handle_authentication_error(
|
||||
original_exception,
|
||||
MagicMock(),
|
||||
{},
|
||||
"/v1/chat/completions",
|
||||
None,
|
||||
"undefined",
|
||||
)
|
||||
finally:
|
||||
verbose_proxy_logger.propagate = False
|
||||
|
||||
records = [r for r in caplog.records if "user_api_key_auth(): Exception occured" in r.getMessage()]
|
||||
assert len(records) == 1
|
||||
assert records[0].name == verbose_proxy_logger.name
|
||||
assert records[0].levelno == logging.ERROR
|
||||
assert records[0].exc_info is not None
|
||||
assert "Authentication service temporarily unavailable" in records[0].getMessage()
|
||||
assert exc_info.value.code == str(status.HTTP_503_SERVICE_UNAVAILABLE)
|
||||
assert not hasattr(exc_info.value, "_litellm_invalid_virtual_key_error")
|
||||
|
|
|
|||
|
|
@ -1,15 +1,15 @@
|
|||
import ast
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
import sys
|
||||
from io import StringIO
|
||||
from pathlib import Path
|
||||
from typing import List
|
||||
|
||||
import pytest
|
||||
|
||||
import logging
|
||||
|
||||
import litellm
|
||||
from litellm._logging import (
|
||||
_COLOR_LOG_FORMAT,
|
||||
|
|
@ -32,6 +32,7 @@ from litellm._logging import (
|
|||
trace_id_var,
|
||||
verbose_logger,
|
||||
verbose_proxy_logger,
|
||||
verbose_proxy_stdout_logger,
|
||||
verbose_router_logger,
|
||||
)
|
||||
from litellm.constants import LITELLM_TRUNCATED_PAYLOAD_FIELD
|
||||
|
|
@ -78,6 +79,21 @@ def test_json_mode_emits_one_record_per_logger(capfd):
|
|||
assert "timestamp" in obj, "`timestamp` key missing"
|
||||
|
||||
|
||||
def test_json_mode_routes_invalid_key_record_once_to_stdout(capfd):
|
||||
_turn_on_json()
|
||||
|
||||
verbose_proxy_stdout_logger.warning("invalid virtual key")
|
||||
|
||||
out, err = capfd.readouterr()
|
||||
assert [raw for raw in err.splitlines() if raw.strip()] == []
|
||||
lines = [raw for raw in out.splitlines() if raw.strip()]
|
||||
assert len(lines) == 1, f"got {len(lines)} lines, want 1: {lines!r}"
|
||||
|
||||
record = json.loads(lines[0])
|
||||
assert record["message"] == "invalid virtual key"
|
||||
assert record["component"] == verbose_proxy_stdout_logger.name
|
||||
|
||||
|
||||
def test_json_formatter_parses_embedded_json_message():
|
||||
"""
|
||||
Test that JsonFormatter parses embedded JSON in the message field and promotes
|
||||
|
|
@ -845,32 +861,121 @@ class _FakeStream:
|
|||
return self._tty
|
||||
|
||||
|
||||
def test_records_below_warning_go_to_stdout_and_the_rest_to_stderr(capsys):
|
||||
logger = logging.getLogger("test_level_routing")
|
||||
def test_invalid_key_warning_routes_as_json_to_stdout(capsys):
|
||||
logger = logging.getLogger("LiteLLM Proxy.stdout")
|
||||
original_handlers = logger.handlers[:]
|
||||
original_propagate = logger.propagate
|
||||
original_level = logger.level
|
||||
logger.handlers.clear()
|
||||
logger.propagate = False
|
||||
logger.setLevel(logging.DEBUG)
|
||||
handler = LevelRoutingStreamHandler()
|
||||
handler.setFormatter(JsonFormatter())
|
||||
logger.addHandler(handler)
|
||||
|
||||
try:
|
||||
logger.warning("invalid virtual key")
|
||||
finally:
|
||||
logger.handlers = original_handlers
|
||||
logger.propagate = original_propagate
|
||||
logger.setLevel(original_level)
|
||||
|
||||
out, err = capsys.readouterr()
|
||||
assert err == ""
|
||||
log_record = json.loads(out)
|
||||
assert log_record["message"] == "invalid virtual key"
|
||||
assert log_record["level"] == "WARNING"
|
||||
|
||||
|
||||
def test_records_below_warning_and_invalid_key_warnings_go_to_stdout(capsys):
|
||||
logger = logging.getLogger("test_level_routing")
|
||||
invalid_key_logger = logging.getLogger("LiteLLM Proxy.stdout")
|
||||
original_handlers = logger.handlers[:]
|
||||
original_propagate = logger.propagate
|
||||
original_level = logger.level
|
||||
original_invalid_key_handlers = invalid_key_logger.handlers[:]
|
||||
original_invalid_key_propagate = invalid_key_logger.propagate
|
||||
original_invalid_key_level = invalid_key_logger.level
|
||||
logger.handlers.clear()
|
||||
logger.propagate = False
|
||||
logger.setLevel(logging.DEBUG)
|
||||
invalid_key_logger.handlers.clear()
|
||||
invalid_key_logger.propagate = False
|
||||
invalid_key_logger.setLevel(logging.DEBUG)
|
||||
handler = LevelRoutingStreamHandler()
|
||||
handler.setFormatter(logging.Formatter("%(levelname)s %(message)s"))
|
||||
logger.addHandler(handler)
|
||||
invalid_key_logger.addHandler(handler)
|
||||
|
||||
try:
|
||||
logger.debug("d")
|
||||
logger.info("i")
|
||||
logger.warning("w")
|
||||
invalid_key_logger.warning("invalid-key-warning")
|
||||
logger.error("e")
|
||||
logger.critical("c")
|
||||
finally:
|
||||
logger.handlers.clear()
|
||||
logger.handlers = original_handlers
|
||||
logger.propagate = original_propagate
|
||||
logger.setLevel(original_level)
|
||||
invalid_key_logger.handlers = original_invalid_key_handlers
|
||||
invalid_key_logger.propagate = original_invalid_key_propagate
|
||||
invalid_key_logger.setLevel(original_invalid_key_level)
|
||||
|
||||
out, err = capsys.readouterr()
|
||||
assert out.splitlines() == ["DEBUG d", "INFO i"]
|
||||
assert out.splitlines() == ["DEBUG d", "INFO i", "WARNING invalid-key-warning"]
|
||||
assert err.splitlines() == ["WARNING w", "ERROR e", "CRITICAL c"]
|
||||
|
||||
|
||||
def test_verbose_loggers_route_records_by_level():
|
||||
for lg in (verbose_logger, verbose_router_logger, verbose_proxy_logger):
|
||||
for lg in (verbose_logger, verbose_router_logger, verbose_proxy_logger, verbose_proxy_stdout_logger):
|
||||
assert any(isinstance(h, LevelRoutingStreamHandler) for h in lg.handlers), lg.name
|
||||
assert verbose_proxy_stdout_logger.level == logging.WARNING
|
||||
assert verbose_proxy_stdout_logger.handlers[0].level == logging.WARNING
|
||||
|
||||
|
||||
def test_invalid_virtual_key_record_does_not_propagate_to_root_handler():
|
||||
root_logger = logging.getLogger()
|
||||
root_stream = StringIO()
|
||||
root_handler = logging.StreamHandler(root_stream)
|
||||
root_handler.setFormatter(logging.Formatter("ROOT %(levelname)s %(message)s"))
|
||||
root_logger.addHandler(root_handler)
|
||||
|
||||
try:
|
||||
verbose_proxy_stdout_logger.warning("invalid virtual key")
|
||||
finally:
|
||||
root_logger.removeHandler(root_handler)
|
||||
|
||||
assert root_stream.getvalue() == ""
|
||||
|
||||
|
||||
def test_turn_on_json_preserves_invalid_key_warning_visibility(monkeypatch, capfd):
|
||||
monkeypatch.setenv("LITELLM_LOG", "ERROR")
|
||||
_turn_on_json()
|
||||
|
||||
verbose_proxy_stdout_logger.warning("invalid virtual key")
|
||||
|
||||
out, err = capfd.readouterr()
|
||||
assert [raw for raw in err.splitlines() if raw.strip()] == []
|
||||
records = [json.loads(raw) for raw in out.splitlines() if raw.strip()]
|
||||
assert len(records) == 1
|
||||
assert records[0]["component"] == verbose_proxy_stdout_logger.name
|
||||
assert records[0]["level"] == "WARNING"
|
||||
|
||||
|
||||
def test_ordinary_proxy_records_still_propagate_to_root_handler():
|
||||
root_logger = logging.getLogger()
|
||||
root_stream = StringIO()
|
||||
root_handler = logging.StreamHandler(root_stream)
|
||||
root_handler.setFormatter(logging.Formatter("ROOT %(levelname)s %(message)s"))
|
||||
root_logger.addHandler(root_handler)
|
||||
|
||||
try:
|
||||
verbose_proxy_logger.error("ordinary proxy error")
|
||||
finally:
|
||||
root_logger.removeHandler(root_handler)
|
||||
|
||||
assert "ROOT ERROR ordinary proxy error" in root_stream.getvalue()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue