mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
fix(passthrough): redact the upstream error preview in original coordinates
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
e4fb4171df
commit
7267f32954
3 changed files with 113 additions and 22 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue