diff --git a/litellm/constants.py b/litellm/constants.py index 137eea15174..4342c7d7ebe 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1532,6 +1532,7 @@ PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS: Final = 4096 PASSTHROUGH_UPSTREAM_ERROR_REPORT_TASK_NAME: Final = "passthrough-upstream-error-report" PASSTHROUGH_UPSTREAM_ERROR_REPORT_SHUTDOWN_WAIT_SECONDS: Final = 10.0 PASSTHROUGH_UPSTREAM_ERROR_REDACT_MARGIN_CHARS: Final = 256 +PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_SCAN_BYTES: Final = 65536 SSE_LATE_RELAY_BACKGROUND_WAIT_SECONDS: Final = 10.0 # Headers to control callbacks diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index b015f5afe71..aae04fa96e2 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -44,6 +44,7 @@ from litellm._uuid import uuid from litellm.constants import ( MAXIMUM_TRACEBACK_LINES_TO_LOG, PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS, + PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_SCAN_BYTES, PASSTHROUGH_UPSTREAM_ERROR_REDACT_MARGIN_CHARS, PASSTHROUGH_UPSTREAM_ERROR_REPORT_TASK_NAME, REDACTED_BY_LITELLM, @@ -863,13 +864,10 @@ def _resolve_team_callback_wiring( ) -def _truncate_upstream_error_body(body: str) -> str: - if len(body) <= PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS: +def _truncate_upstream_error_body(body: str, *, truncated: bool) -> str: + if not truncated: return body - return ( - f"{body[:PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS]}... " - f"(truncated at {PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS} chars)" - ) + return f"{body}... (truncated at {PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS} chars)" def _sanitize_upstream_error_body(body: str) -> str: @@ -949,7 +947,7 @@ def _headers_without_body_framing(headers: httpx.Headers) -> httpx.Headers: ) -_DANGLING_PRIVATE_KEY: Final = re.compile(r"-----BEGIN[A-Z \-]*PRIVATE KEY-----") +_DANGLING_PRIVATE_KEY: Final = re.compile(r"-----BEGIN[A-Z \-]*PRIVATE KEY----- ?[A-Za-z0-9+/= ]{32,}") def _mask_dangling_private_key(text: str) -> str: @@ -957,6 +955,29 @@ def _mask_dangling_private_key(text: str) -> str: return text if match is None else text[: match.start()] + REDACTED +def _redact_preview_head(text: str, redact: Callable[[str], str]) -> str: + head: Final = text[:PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS] + margin: Final = text[ + PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS : PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS + + PASSTHROUGH_UPSTREAM_ERROR_REDACT_MARGIN_CHARS + ] + redacted_head: Final = redact(head) + if not margin or redact(head + margin) == redacted_head + redact(margin): + return redacted_head + return redacted_head[: redacted_head.rfind(" ") + 1] + REDACTED + + +def _upstream_error_body_for_log(preview: bytes, encoding: str, redact: Callable[[str], str]) -> str: + sanitized: Final = _sanitize_upstream_error_body( + preview[:PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_SCAN_BYTES].decode(encoding, errors="replace") + ) + return _truncate_upstream_error_body( + _mask_dangling_private_key(_redact_preview_head(sanitized, redact)), + truncated=len(preview) > PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_SCAN_BYTES + or len(sanitized) > PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS, + ) + + def _passthrough_upstream_failure_reporter( response: httpx.Response, user_api_key_dict: UserAPIKeyAuth, @@ -967,17 +988,10 @@ def _passthrough_upstream_failure_reporter( redact: Callable[[str], str] = redact_secrets, ) -> _ReportPreview: async def report(preview: bytes) -> None: - preview_text: Final = preview[ - : (PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS + PASSTHROUGH_UPSTREAM_ERROR_REDACT_MARGIN_CHARS) * 4 - ].decode(response.encoding or "utf-8", errors="replace")[ - : PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS + PASSTHROUGH_UPSTREAM_ERROR_REDACT_MARGIN_CHARS - ] upstream_error_body: Final = ( REDACTED_BY_LITELLM if should_redact_message_logging(logging_obj.model_call_details) - else _truncate_upstream_error_body( - _sanitize_upstream_error_body(_mask_dangling_private_key(redact(preview_text))) - ) + else _upstream_error_body_for_log(preview, response.encoding or "utf-8", redact) ) log_warning( "pass_through_endpoint: upstream %s %s returned %s: %s", diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 6f444c7db04..7ab281e7192 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -23,8 +23,8 @@ from starlette.datastructures import FormData, Headers, QueryParams from starlette.datastructures import UploadFile as StarletteUploadFile import litellm -from litellm._logging import verbose_proxy_logger -from litellm.constants import DEFAULT_REQUEST_TIMEOUT_SECONDS +from litellm._logging import redact_secrets, verbose_proxy_logger +from litellm.constants import DEFAULT_REQUEST_TIMEOUT_SECONDS, PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.proxy._types import ProxyException, UserAPIKeyAuth @@ -35,6 +35,8 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( InitPassThroughEndpointHelpers, _PreviewReportingStream, _passthrough_upstream_failure_reporter, + _sanitize_upstream_error_body, + _upstream_error_body_for_log, _registered_pass_through_routes, _truncate_upstream_error_body, _with_trace_context, @@ -4297,11 +4299,12 @@ async def test_pass_through_request_streaming_upstream_error_body_reaches_client @pytest.mark.asyncio async def test_truncate_upstream_error_body_caps_at_log_limit(): short_body: Final = "x" * 4096 - assert _truncate_upstream_error_body(short_body) == short_body + assert _truncate_upstream_error_body(short_body, truncated=False) == short_body + + truncated: Final = _truncate_upstream_error_body("a" * 4096, truncated=True) + assert truncated == f"{'a' * 4096}... (truncated at 4096 chars)" long_body: Final = "a" * 5000 - truncated: Final = _truncate_upstream_error_body(long_body) - assert truncated == f"{'a' * 4096}... (truncated at 4096 chars)" upstream_response: Final = httpx.Response( status_code=500, @@ -8326,8 +8329,8 @@ async def test_passthrough_upstream_failure_reporter_redacts_only_the_bounded_pr ) await report(preview) max_plus_margin: Final = 4096 + 256 - assert len(redacted_inputs) == 1, redacted_inputs - assert len(redacted_inputs[0]) <= max_plus_margin, len(redacted_inputs[0]) + assert 1 <= len(redacted_inputs) <= 3, redacted_inputs + assert max(len(text) for text in redacted_inputs) <= max_plus_margin, redacted_inputs assert len(logged_details) == 1, logged_details assert "REDACTED-KEY" in logged_details[0], logged_details[0] assert marker_key not in logged_details[0], logged_details[0] @@ -8416,3 +8419,76 @@ async def test_passthrough_upstream_failure_reporter_keeps_text_after_a_complete assert len(logged_details) == 1, logged_details assert "MIIEvwIBADANBgkqhkiG9w0BAQEFAASCBKkwggSlAgEAAoIBAQshortkey==" not in logged_details[0], logged_details[0] assert "INTERNAL" in logged_details[0], logged_details[0] + + +_TRUNCATION_MARKER: Final = f"... (truncated at {PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS} chars)" + + +@pytest.mark.asyncio +async def test_upstream_error_body_for_log_keeps_a_bare_pem_header_unmasked(): + """Vertex-style error text mentions the PEM header literally; only a header + followed by real key material may be masked.""" + preview: Final = ( + '{"error":"private_key field must start with -----BEGIN PRIVATE KEY----- header, got garbage; request_id=abc"}' + ).encode() + output: Final = _upstream_error_body_for_log(preview, "utf-8", redact_secrets) + assert "-----BEGIN PRIVATE KEY----- header" in output, output + assert "request_id=abc" in output, output + assert "REDACTED" not in output, output + + +@pytest.mark.asyncio +async def test_upstream_error_body_for_log_drops_a_secret_fragment_cut_at_the_preview_edge(): + """A secret whose match spans the head/margin cut must not leak its head into + the log: the straddle branch drops the trailing partial and writes REDACTED.""" + aws_key: Final = "AKIAIOSFODNN7EXAMPLE" + pem: Final = "-----BEGIN PRIVATE KEY----- " + "b64chars" * 75 + " -----END PRIVATE KEY-----" + pad_len: Final = 4096 + 256 - 19 - len(pem) - 2 + assert pad_len > 0 + sanitized: Final = _sanitize_upstream_error_body(pem + " " + "x" * pad_len + " " + aws_key) + assert sanitized.index(aws_key) == 4096 + 256 - 19 + output: Final = _upstream_error_body_for_log(sanitized.encode(), "utf-8", redact_secrets) + assert "AKIA" not in output, output + assert "AKIAIOSFODNN7EXAMPL" not in output, output + assert "REDACTED" in output, output + assert output.endswith(_TRUNCATION_MARKER), output + + +@pytest.mark.asyncio +async def test_upstream_error_body_for_log_keeps_non_secret_straddle_text(): + """Plain text that simply runs past the head must survive unchanged: the + straddle branch only fires when a redaction match spans the cut.""" + sanitized: Final = "word " * 840 + preview: Final = sanitized.encode() + output: Final = _upstream_error_body_for_log(preview, "utf-8", redact_secrets) + expected: Final = sanitized[:PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS] + _TRUNCATION_MARKER + assert output == expected, output[:200] + + +@pytest.mark.asyncio +async def test_upstream_error_body_for_log_sanitizes_before_slicing(): + """Whitespace must collapse before the head is taken, or the useful part of a + pretty-printed error body would be mostly indentation.""" + body: Final = json.dumps( + {"error": {"message": "x" * 100, "details": [{"k": f"v{i}", "pad": " " * 40} for i in range(400)]}}, + indent=4, + ) + sanitized: Final = _sanitize_upstream_error_body(body) + assert len(sanitized) > PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS, len(sanitized) + output: Final = _upstream_error_body_for_log(body.encode(), "utf-8", redact_secrets) + useful: Final = output.removesuffix(_TRUNCATION_MARKER) + assert len(useful) == PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS, len(useful) + assert output.endswith(_TRUNCATION_MARKER), output + assert useful == sanitized[:PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS], useful[:200] + + +@pytest.mark.asyncio +async def test_upstream_error_body_for_log_marks_bodies_past_the_scan_cap(): + """A body beyond the 64 KiB raw scan cap must carry the truncation marker + even when its scanned head collapses below the log cap: the tail we never + decoded is unaccounted for.""" + preview: Final = b"a" * 3000 + b" " * 5_000_000 + output: Final = _upstream_error_body_for_log(preview, "utf-8", redact_secrets) + useful: Final = output.removesuffix(_TRUNCATION_MARKER) + assert useful == "a" * 3000, useful[:200] + assert output.endswith(_TRUNCATION_MARKER), output