Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/openai_files_endpoints/batch_guardrails.py: 32%
254 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""
2Run the configured pre-call guardrails over every record of a batch input file.
4Runs after ``batch_file_validation.check_batch_file_upload``, so every line here is already known
5to parse as a JSON object carrying ``custom_id``, ``method``, ``url`` and ``body``.
6"""
8from __future__ import annotations
10import asyncio
11import copy
12import json
13import re
14import tempfile
15from collections.abc import Iterator, Mapping
16from dataclasses import dataclass
17from types import MappingProxyType
18from typing import TYPE_CHECKING, BinaryIO, Final, NoReturn, TypeAlias
19from urllib.parse import urlsplit
21from fastapi import HTTPException
22from typing_extensions import assert_never
24from litellm.exceptions import GuardrailRaisedException
25from litellm.integrations.custom_guardrail import is_guardrail_intervention
26from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route
27from litellm.proxy._types import UserAPIKeyAuth
28from litellm.types.llms.openai import BatchGuardrailRecord, BatchGuardrailReport
29from litellm.types.utils import CallTypes, CallTypesLiteral
31if TYPE_CHECKING: 31 ↛ 32line 31 didn't jump to line 32 because the condition on line 31 was never true
32 from litellm.proxy.utils import ProxyLogging
34EMPTY_MAPPING: Final[Mapping[str, object]] = MappingProxyType({})
36_SCAN_WINDOW: Final = 32
38# Past this the rewrite rolls to disk, keeping the router's per-deployment deepcopy of the handle
39# as cheap as it is for the spooled upload this replaces.
40_REWRITE_SPOOL_BYTES: Final = 1024 * 1024
42# custom_id is caller-supplied and reaches a log line, so it is stripped of control characters
43# and capped rather than rendered as given.
44_CONTROL_CHARACTERS: Final = re.compile(r"[\x00-\x1f\x7f]")
45_CUSTOM_ID_LOG_LIMIT: Final = 128
46_SUMMARY_LIMIT: Final = 50
48_SCAN_METADATA_KEY: Final = "litellm_metadata"
49_SCAN_METADATA_BAGS: Final = (_SCAN_METADATA_KEY, "metadata")
51# Set by pre_call_hook when a guardrail rerouted the request to a different model.
52_ROUTE_APPLIED_KEY: Final = "sensitive_data_routing_applied"
54# Dropped before dispatch and restored afterwards rather than diffed. Guardrail dispatch writes
55# its bookkeeping into `metadata`, and a record's own metadata is not scanned content on the
56# online path either. `guardrails` is dropped because guardrail selection reads it ahead of the
57# proxy-injected list, so leaving it would let a record's own body opt out of the chain its key
58# and team selected; online that key can only add to the list, never replace it.
59_INJECTED_KEYS: Final = frozenset({_SCAN_METADATA_KEY, "metadata", "guardrails"})
61# Only what guardrail dispatch reads. The parent OTel span is deliberately left out: parenting one
62# guardrail span per record would put tens of thousands of spans on a single upload's trace.
63_SCAN_METADATA_KEYS: Final = frozenset(
64 {
65 "guardrails",
66 "_guardrail_pipelines",
67 "_pipeline_managed_guardrails",
68 "user_api_key_metadata",
69 "user_api_key_team_metadata",
70 "tags",
71 "headers",
72 }
73)
75_SCANNABLE_CALL_TYPES: Final = frozenset(
76 {
77 CallTypes.acompletion,
78 CallTypes.atext_completion,
79 CallTypes.aembedding,
80 CallTypes.aresponses,
81 CallTypes.anthropic_messages,
82 }
83)
85# Mirrors the record classifier in litellm/llms/bedrock/files/transformation.py, so a record
86# litellm already accepts without a url keeps working.
87_BODY_SHAPE_CALL_TYPES: Final = (
88 ("messages", CallTypes.acompletion),
89 ("prompt", CallTypes.atext_completion),
90 ("input", CallTypes.aembedding),
91)
94@dataclass(frozen=True, slots=True)
95class UnparseableRecord:
96 line_number: int
99@dataclass(frozen=True, slots=True)
100class UnscannableRecord:
101 line_number: int
102 custom_id: str | None
103 url: str | None
106@dataclass(frozen=True, slots=True)
107class UnroutableRecord:
108 line_number: int
109 custom_id: str | None
110 guardrail: str | None
113BatchScanFailure: TypeAlias = UnparseableRecord | UnscannableRecord | UnroutableRecord
116@dataclass(frozen=True, slots=True)
117class _Redaction:
118 """A rewritten record on its way to the scan spool, held only for the window it was scanned in."""
120 line_number: int
121 custom_id: str | None
122 text: str
125@dataclass(frozen=True, slots=True)
126class RecordRedacted:
127 line_number: int
128 custom_id: str | None
129 offset: int
130 length: int
131 """Where the re-serialized record sits in the scan spool, so a large file's rewrites stay off the heap."""
134@dataclass(frozen=True, slots=True)
135class RecordDropped:
136 line_number: int
137 custom_id: str | None
138 guardrail: str | None = None
141_RecordChange: TypeAlias = RecordRedacted | RecordDropped
142_ScanOutcome: TypeAlias = BatchScanFailure | _Redaction | RecordDropped
145@dataclass(frozen=True, slots=True)
146class BatchScanResult:
147 """What the scan decided, per record. Empty changes means the upload proceeds untouched."""
149 changes: tuple[_RecordChange, ...]
150 scanned_records: int
151 redactions: BinaryIO
152 """Spool holding every rewritten record, keyed by the offsets on each ``RecordRedacted``."""
154 @property
155 def submitted_records(self) -> int:
156 return self.scanned_records - sum(1 for change in self.changes if isinstance(change, RecordDropped))
158 def summary(self) -> str:
159 """Compact per-record outcome for the server-side log line, capped so one upload cannot flood it."""
160 shown: Final = ", ".join(
161 f"line {change.line_number}{_describe(change.custom_id)} "
162 f"{'redacted' if isinstance(change, RecordRedacted) else 'dropped'}"
163 for change in self.changes[:_SUMMARY_LIMIT]
164 )
165 remaining: Final = len(self.changes) - _SUMMARY_LIMIT
166 return shown if remaining <= 0 else f"{shown}, and {remaining} more"
168 def report(self) -> BatchGuardrailReport:
169 return BatchGuardrailReport(
170 submitted_records=self.submitted_records,
171 modified_records=tuple(
172 BatchGuardrailRecord(
173 line=change.line_number,
174 custom_id=change.custom_id,
175 action="redacted" if isinstance(change, RecordRedacted) else "dropped",
176 guardrail=change.guardrail if isinstance(change, RecordDropped) else None,
177 )
178 for change in self.changes
179 ),
180 )
183@dataclass(frozen=True, slots=True)
184class _ParsedRecord:
185 line_number: int
186 payload: Mapping[str, object]
189def _rejected(message: str) -> HTTPException:
190 return HTTPException(status_code=400, detail={"error": message}) # mutable-ok: FastAPI detail shape
193def raise_public(failure: BatchScanFailure) -> NoReturn:
194 """Map a scan failure onto the 400 contract the files endpoint already returns."""
195 match failure:
196 case UnparseableRecord(line_number=line_number):
197 raise _rejected(
198 f"The 'body' of batch input line {line_number} is not an object, so guardrails cannot be applied to it"
199 )
200 case UnscannableRecord(line_number=line_number, custom_id=custom_id, url=url):
201 raise _rejected(
202 f"Batch input line {line_number}{_describe(custom_id)} targets {url or 'no url'} "
203 "and its body has no messages, prompt or input, so guardrails cannot read it. "
204 "Give the record a chat, completion, embedding, responses or messages body"
205 )
206 case UnroutableRecord(line_number=line_number, custom_id=custom_id, guardrail=guardrail):
207 raise _rejected(
208 f"Batch input line {line_number}{_describe(custom_id)} was routed to a different model by "
209 f"{guardrail or 'a guardrail'}, and every record of a batch file goes to one provider, so "
210 "the file cannot be submitted. Send that record outside the batch"
211 )
212 case _:
213 assert_never(failure)
216def raise_nothing_to_submit() -> NoReturn:
217 """Every record was blocked, so there is no batch left to create."""
218 raise _rejected(
219 "Every record in the batch input file was blocked by a guardrail, so there is nothing left to submit"
220 )
223def _is_content_block(exc: BaseException) -> bool:
224 """
225 Whether the guardrail judged the record, as opposed to failing to judge it.
227 Stricter than ``is_guardrail_intervention``, which answers a different question and counts
228 every ``GuardrailRaisedException`` as a block. Several integrations raise that same exception
229 for an unreachable backend or an unparseable response, and only when the operator configured
230 the guardrail to fail closed, so treating it as a block would turn "refuse this request" into
231 "drop this record and submit the rest", which is the silent loss of enforcement this whole
232 path exists to prevent. A guardrail that does not say it blocked content aborts the upload.
234 Guardrails that report a technical failure as an ``HTTPException`` carrying a block status
235 are caught by ``__cause__``: raising ``from`` the underlying error is a deliberate statement
236 that something else caused this, which a verdict on content never is. Implicit context is
237 left alone, since a block raised inside an unrelated ``except`` would read as a failure.
238 """
239 if isinstance(exc, GuardrailRaisedException):
240 return exc.blocked_content
241 if exc.__cause__ is not None:
242 return False
243 return is_guardrail_intervention(exc)
246def _naming_guardrail(exc: BaseException) -> str | None:
247 """The guardrail that raised, from whichever place it recorded its own name."""
248 named: Final = getattr(exc, "guardrail_name", None)
249 if isinstance(named, str):
250 return named
251 detail: Final = getattr(exc, "detail", None)
252 enriched: Final = detail.get("guardrail_name") if isinstance(detail, dict) else None
253 return enriched if isinstance(enriched, str) else None
256def _describe(custom_id: str | None) -> str:
257 if not custom_id:
258 return ""
259 safe: Final = _CONTROL_CHARACTERS.sub(" ", custom_id)[:_CUSTOM_ID_LOG_LIMIT]
260 return f" (custom_id {safe})"
263def _iter_lines(source: BinaryIO) -> Iterator[tuple[int, bytes]]:
264 """
265 Yield every non-blank line with its 1-based number, so both passes number records alike.
267 Bytes, not text. The upload validation immediately before this parses each line as bytes,
268 where the json module sniffs the encoding itself and accepts a leading byte order mark or a
269 lone surrogate. Decoding to `str` first is stricter than that, so a file written by any of
270 the editors that emit a BOM would pass validation and then fail the scan.
271 """
272 for line_number, raw_line in enumerate(source, start=1):
273 if raw_line.strip():
274 yield line_number, raw_line
277def _iter_records(source: BinaryIO) -> Iterator[_ParsedRecord]:
278 """Yield one record per line, relying on the upload validation that already ran."""
279 for line_number, raw_line in _iter_lines(source):
280 yield _ParsedRecord(line_number=line_number, payload=json.loads(raw_line))
283def _call_type_from_url(url: str) -> CallTypesLiteral | None:
284 """
285 Resolve the route a record names, tolerating how callers actually write it.
287 An absolute url has to reduce to its path or nothing matches, and a record naming
288 ``/v1/responses`` in full would fall through to its body, where ``input`` reads as an
289 embedding and the record gets scanned as the wrong call type rather than the right one.
290 """
291 try:
292 path: Final = urlsplit(url).path.split("?")[0].rstrip("/")
293 except ValueError:
294 # urlsplit rejects a few malformed authorities outright, and the validation that ran
295 # before this only checks the key is present. An unreadable url is one we do not
296 # recognize, which is what falling back to the body shape already handles.
297 return None
298 call_types: Final = get_call_types_for_route(path)
299 if call_types is None:
300 return None
301 scannable: Final = next((c for c in call_types if c in _SCANNABLE_CALL_TYPES), None)
302 return None if scannable is None else scannable.value
305def _call_type_from_body(body: Mapping[str, object]) -> CallTypesLiteral | None:
306 shape: Final = next((call_type for field, call_type in _BODY_SHAPE_CALL_TYPES if field in body), None)
307 return None if shape is None else shape.value
310def _scannable_call_type(url: object, body: Mapping[str, object]) -> CallTypesLiteral | None:
311 """
312 Resolve how to scan a record: its url when we recognize one, otherwise its body shape.
314 An unrecognized url falls through to the body rather than rejecting, because a record we can
315 still read is a record we can still scan, and the provider transformers treat an unknown url
316 as chat rather than as an error.
317 """
318 from_url: Final = _call_type_from_url(url) if isinstance(url, str) and url else None
319 return from_url if from_url is not None else _call_type_from_body(body)
322def _custom_id_of(payload: Mapping[str, object]) -> str | None:
323 """
324 The record's identifier, rendered as text.
326 The batch spec asks for a string, but callers do send numbers, and reporting those as null
327 would leave the one field a caller reconciles on empty for exactly the records it needs.
328 """
329 custom_id: Final = payload.get("custom_id")
330 if isinstance(custom_id, str):
331 # A lone surrogate parses out of the file but cannot be encoded back out, and this value
332 # is echoed in the response, so rendering it would fail the whole upload with a 500.
333 return custom_id.encode("utf-8", "replace").decode("utf-8")
334 return str(custom_id) if isinstance(custom_id, (int, float)) and not isinstance(custom_id, bool) else None
337def _fingerprint(body: Mapping[str, object], keys: frozenset[str]) -> str:
338 """
339 Order-insensitive projection, so a guardrail re-serializing a dict does not read as a change.
341 An absent key projects to ``null`` while a key holding ``None`` projects to the string
342 ``"null"``, so adding or dropping a null-valued key still reads as a change.
343 """
344 return json.dumps(
345 tuple(
346 (key, json.dumps(body[key], sort_keys=True, default=str) if key in body else None) for key in sorted(keys)
347 )
348 )
351def build_scan_metadata(request_metadata: Mapping[str, object]) -> Mapping[str, object]:
352 """
353 Narrow the request metadata to the keys guardrail dispatch reads.
355 Passing the whole thing through would carry values that cannot be copied, such as the parent
356 OTel span, and would hand every record proxy state it has no business seeing.
357 """
358 return MappingProxyType({key: value for key, value in request_metadata.items() if key in _SCAN_METADATA_KEYS})
361async def _scan_record(
362 record: _ParsedRecord,
363 scan_metadata: Mapping[str, object],
364 user_api_key_dict: UserAPIKeyAuth,
365 proxy_logging_obj: ProxyLogging,
366) -> _ScanOutcome | None:
367 body: Final = record.payload.get("body")
368 if not isinstance(body, dict):
369 return UnparseableRecord(line_number=record.line_number)
371 custom_id: Final = _custom_id_of(record.payload)
372 url: Final = record.payload.get("url")
373 call_type: Final = _scannable_call_type(url, body)
374 if call_type is None:
375 return UnscannableRecord(
376 line_number=record.line_number,
377 custom_id=custom_id,
378 url=url if isinstance(url, str) else None,
379 )
381 scan_input: Final[dict[str, object]] = copy.deepcopy(body) # mutable-ok: pre_call_hook mutates the dict it is given
382 own_injected: Final = MappingProxyType({key: body[key] for key in _INJECTED_KEYS if key in body})
383 for injected in _INJECTED_KEYS:
384 scan_input.pop(injected, None)
385 # Both bags, because guardrails read whichever one their own route populates and a record
386 # scanned as chat reaches ones that only ever look at `metadata`; both are injected keys, so
387 # neither survives into the record that ships. Deep, and per bag per record, because `headers`
388 # and `tags` are nested containers otherwise shared with the upload request and with every
389 # other record in the window. The narrowing above already removed what cannot be copied.
390 for injected in _SCAN_METADATA_BAGS:
391 scan_input[injected] = copy.deepcopy(dict(scan_metadata)) # mutable-ok: guardrails write here
393 try:
394 # The chain hands back the body it produced, which may be a replacement for the dict it was
395 # given rather than that same dict mutated, so this is what gets compared.
396 scanned: Final[dict] = await proxy_logging_obj.pre_call_hook( # mutable-ok: the guardrails' own dict
397 user_api_key_dict=user_api_key_dict,
398 data=scan_input,
399 call_type=call_type,
400 guardrails_only=True,
401 )
402 except Exception as exc:
403 if _is_content_block(exc):
404 return RecordDropped(line_number=record.line_number, custom_id=custom_id, guardrail=_naming_guardrail(exc))
405 raise
407 rerouted: Final = scanned.get("metadata")
408 if isinstance(rerouted, dict) and rerouted.get(_ROUTE_APPLIED_KEY):
409 return UnroutableRecord(
410 line_number=record.line_number,
411 custom_id=custom_id,
412 guardrail=rerouted.get("sensitive_data_routing_guardrail"),
413 )
415 compared: Final = (frozenset(body) | frozenset(scanned)) - _INJECTED_KEYS
416 if _fingerprint(scanned, compared) == _fingerprint(body, compared):
417 return None
418 for injected in _INJECTED_KEYS:
419 scanned.pop(injected, None)
420 scanned.update(own_injected)
421 return _Redaction(
422 line_number=record.line_number,
423 custom_id=custom_id,
424 text=json.dumps({**record.payload, "body": scanned}), # mutable-ok: json.dumps needs a plain dict
425 )
428async def _scan_window(
429 window: tuple[_ParsedRecord, ...],
430 scan_metadata: Mapping[str, object],
431 user_api_key_dict: UserAPIKeyAuth,
432 proxy_logging_obj: ProxyLogging,
433) -> tuple[tuple[int, _ScanOutcome | BaseException], ...]:
434 """``return_exceptions=True`` so one record raising never leaves its siblings unobserved."""
435 outcomes: Final = await asyncio.gather(
436 *(_scan_record(record, scan_metadata, user_api_key_dict, proxy_logging_obj) for record in window),
437 return_exceptions=True,
438 )
439 return tuple((record.line_number, outcome) for record, outcome in zip(window, outcomes) if outcome is not None)
442def _spool(redactions: BinaryIO, redaction: _Redaction) -> RecordRedacted:
443 """Park the rewritten record on disk so only its location is carried for the rest of the scan."""
444 encoded: Final = redaction.text.encode("utf-8")
445 redactions.seek(0, 2)
446 offset: Final = redactions.tell()
447 redactions.write(encoded)
448 return RecordRedacted(
449 line_number=redaction.line_number,
450 custom_id=redaction.custom_id,
451 offset=offset,
452 length=len(encoded),
453 )
456def _worst(problems: tuple[tuple[int, BatchScanFailure | BaseException], ...]) -> BatchScanFailure | BaseException:
457 """A guardrail that blocked outranks a record we merely refused; then earliest line wins."""
458 raised: Final = tuple(problem for problem in problems if isinstance(problem[1], BaseException))
459 return min(raised or problems, key=lambda problem: problem[0])[1]
462async def scan_batch_input_file(
463 *,
464 file_source: BinaryIO,
465 request_metadata: Mapping[str, object],
466 user_api_key_dict: UserAPIKeyAuth,
467 proxy_logging_obj: ProxyLogging,
468) -> BatchScanFailure | BatchScanResult:
469 """
470 Stream a batch input file and run the pre-call guardrail chain against every record.
472 A record a guardrail rewrites is kept in its rewritten form and a record it blocks is dropped,
473 which is what the online path does per request. Both are returned for reporting. A guardrail
474 exception that is not a block is re-raised untouched so its status code survives, since dropping
475 a record that was never inspected is worse than refusing the file.
476 """
477 scan_metadata: Final = build_scan_metadata(request_metadata)
478 problems: Final[list[tuple[int, BatchScanFailure | BaseException]]] = [] # mutable-ok: spans windows
479 changes: Final[list[_RecordChange]] = [] # mutable-ok: accumulates across windows
480 window: Final[list[_ParsedRecord]] = [] # mutable-ok: bounded read-ahead buffer
481 scanned: Final[list[int]] = [] # mutable-ok: counts records the scan actually reached
482 redactions: Final = tempfile.SpooledTemporaryFile( # noqa: SIM115 # the rewrite reads this back
483 max_size=_REWRITE_SPOOL_BYTES
484 )
486 async def drain() -> None:
487 if window:
488 scanned.append(len(window))
489 for line_number, outcome in await _scan_window(
490 tuple(window), scan_metadata, user_api_key_dict, proxy_logging_obj
491 ):
492 if isinstance(outcome, _Redaction):
493 changes.append(_spool(redactions, outcome))
494 elif isinstance(outcome, RecordDropped):
495 changes.append(outcome)
496 else:
497 problems.append((line_number, outcome))
498 window.clear()
500 try:
501 for item in _iter_records(file_source):
502 window.append(item)
503 if len(window) >= _SCAN_WINDOW:
504 await drain()
505 if problems:
506 break
507 if not problems:
508 await drain()
509 except BaseException:
510 redactions.close()
511 raise
512 finally:
513 file_source.seek(0)
515 if problems:
516 redactions.close()
517 worst: Final = _worst(tuple(problems))
518 if isinstance(worst, BaseException):
519 raise worst
520 return worst
521 if not changes:
522 redactions.close()
523 return BatchScanResult(
524 changes=tuple(sorted(changes, key=lambda change: change.line_number)),
525 scanned_records=sum(scanned),
526 redactions=redactions,
527 )
530def _read_spooled(redactions: BinaryIO, change: RecordRedacted) -> bytes:
531 redactions.seek(change.offset)
532 return redactions.read(change.length)
535def rewrite_batch_input_file(file_source: BinaryIO, result: BatchScanResult) -> BinaryIO:
536 """
537 Re-emit the file with redacted records rewritten and dropped records left out.
539 Untouched records are copied through as written rather than re-serialized, so enabling the
540 feature does not reformat records no guardrail objected to. Blank lines between records are
541 not carried over, since they are not records. Rewritten records are read back from the scan's
542 spool rather than from memory, so a file whose records are mostly rewritten does not put a
543 second copy of itself on the heap.
544 """
545 redacted: Final = MappingProxyType(
546 {change.line_number: change for change in result.changes if isinstance(change, RecordRedacted)}
547 )
548 dropped: Final = frozenset(change.line_number for change in result.changes if isinstance(change, RecordDropped))
550 output: Final = tempfile.SpooledTemporaryFile( # noqa: SIM115 # the caller uploads this handle
551 max_size=_REWRITE_SPOOL_BYTES
552 )
553 wrote_any = False # rebind-ok: tracks whether a separator is needed
554 try:
555 for line_number, raw_line in _iter_lines(file_source):
556 if line_number in dropped:
557 continue
558 change = redacted.get(line_number)
559 line = raw_line.rstrip(b"\n") if change is None else _read_spooled(result.redactions, change)
560 output.write(b"\n" + line if wrote_any else line)
561 wrote_any = True
562 except BaseException:
563 output.close()
564 raise
565 finally:
566 file_source.seek(0)
567 output.seek(0)
568 return output