ci: fail on new unbounded SQL IN lists and add a Prisma chunking helper (#42629)

* ci: warn on SQL IN lists with no written bound

Postgres caps a prepared statement at 32,767 bind parameters and a
membership filter binds one per value, so an IN list built from table
data breaks once the table outgrows the cap. That is how the budget reset
job froze every due budget (LIT-7535, #40564).

check_unbounded_in_lists.py reports every Prisma "in" / "not_in" filter
whose value has no fixed size and every raw SQL literal that splices a
list in after "IN (", unless the line carries "# bounded-ok: <reason>".
It only warns for now: the output is the inventory for RCA action item
AI-1, and it exits 0.

* ci: decide a constant IN list by its module binding, not its casing

An ALL_CAPS name imported or filled at runtime is as unbounded as any
other, so a name now passes only when the module binds it once to a
value of fixed size. Adds Final to the locals a loop does not forbid.

* ci: only a frozen module value makes an IN list constant

A module list bound once could still grow through append or extend, so
a name now counts as fixed only when it is bound to a tuple, frozenset
or constant. Trims the module docstring to what a reader needs.

* ci: chunk Prisma IN lists with a shared helper and fail on new unbounded ones

Add litellm.repositories.bounded_in: find_many_in, count_in, update_many_in
and delete_many_in split a deduplicated value list into 5,000-value chunks,
AND each chunk with the caller's where, run them in order (a transaction
handle works) and combine the results. Writes take a required atomicity
argument, and a where that already filters the chunked field is refused.

check_unbounded_in_lists.py now fails CI on any finding missing from
unbounded_in_baseline.txt and on any stale baseline entry, so the baseline
only shrinks. Entries are keyed by path, enclosing scope, kind, field and
occurrence, not line numbers. The helper module is exempt, a constant
spread into a frozen tuple counts as fixed, and messages point at the
helper for "in" and at an array parameter for "not_in" and raw SQL.

A real-Postgres integration test shows a raw 40,000-value filter rejected
for too many bind variables while the helpers handle it.

* refactor: rename bounded_in to chunked_in and let callers pick a chunk size

The helper module is litellm.repositories.chunked_in, and its unit and
integration tests, the checker's exemption path and its finding messages
follow the new name. The `# bounded-ok` marker is unchanged.

find_many_in, count_in, update_many_in and delete_many_in take a
keyword-only chunk_size, defaulting to IN_LIST_CHUNK_SIZE (5,000). A value
below 1 or above MAX_IN_LIST_CHUNK_SIZE (30,000) raises ValueError before
any query, which leaves the rest of the filter headroom under Postgres's
32,767 bind-parameter cap.

* refactor: flatten chunked_in's stacked comprehensions with chain.from_iterable

LIT014 (#42650) caps a comprehension at one for and one if clause. The four nested walks in the helper now chain their iterables instead, with the same order and results.

* refactor: recover user details with find_many_in, sending chunks as lists

_details_for_user_ids reads users through find_many_in instead of a raw
"in" filter, so its lookup stays under the bind-parameter cap for any
number of recovered keys. Up to 5,000 ids it still sends one find_many
with the same where dict, and a PrismaError from any chunk is still
logged and treated as no details.

The helper now sends each chunk as a list, so a chunked filter equals
the dict a hand-written call would send and a migrated call site's
existing assertions keep passing.

The site's baseline entry is gone.

* ci: skip functional TypedDict field maps in the unbounded IN list check

The dict passed as the field map of TypedDict("Name", {...}), or as its fields= keyword, names fields: an "in" or "notIn" key there is a type, not a filter. Only that dict is skipped, for TypedDict, typing.TypedDict and typing_extensions.TypedDict; a filter nested in a field value or passed to any other call is still reported. The two types/proxy/management_endpoints/team_endpoints.py entries leave the baseline, which is now 156.

* fix: refuse an update_many_in whose data writes the chunked field

Chunks run one after another, so an update that sets the chunked field can move a row into a later chunk, which updates it again and counts it twice: values ["old", "new"] with chunk_size=1 and data={"id": "new"} does exactly that. update_many_in now raises ChunkedFieldWriteError before any query when data has the chunked field as a top-level key, in any form, including Prisma operators such as {"set": ...}.

* docs: cut the unbounded IN list checker's docstring to what it flags and how to clear it

It now says what is reported, the three ways to clear a finding, and how the baseline and --update-baseline work, in 11 lines. The per-shape detail lives in the tests.

* ci: key an unbounded IN list finding by its filtered expression too

A baseline key of path, scope, kind, field and occurrence let a PR delete
a baselined filter and add a different unbounded one on the same field in
the same function, and the new one took over the old key. The key now
also carries the filtered expression's source, whitespace-normalized
(the Prisma value, or a raw-SQL `IN (...)` slot), so that swap reads as
one new and one stale entry and fails the run. The same expression
re-added in the same function is still the same finding.

Every baseline entry is rewritten in the new form; the 156 findings are
unchanged, and only occurrence indexes renumber where one field had
several different expressions.
This commit is contained in:
ryan-crabbe-berri 2026-09-26 13:40:44 -07:00 • committed by GitHub
parent 40297e6268
commit 96c008f420
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
10 changed files with 1715 additions and 3 deletions

View file

@ -146,6 +146,9 @@ jobs:
- name: check_migrations_no_data_rewrites
run: uv run --no-sync python ./tests/code_coverage_tests/check_migrations_no_data_rewrites.py
- name: check_unbounded_in_lists (fails on findings not in the baseline)
run: uv run --no-sync python ./tests/code_coverage_tests/check_unbounded_in_lists.py
- name: memory_test
run: uv run --no-sync python ./tests/code_coverage_tests/memory_test.py

View file

@ -20,6 +20,7 @@ from litellm.constants import (
)
from litellm.litellm_core_utils.litellm_logging import is_valid_sha256_hash
from litellm.proxy.utils import PrismaClient
from litellm.repositories.chunked_in import find_many_in
from litellm.repositories.user_repository import UserRepository
_T = TypeVar("_T")
@ -167,9 +168,7 @@ async def _details_for_user_ids(
if not user_ids:
return _EMPTY_USER_DETAILS
users: Final = await _db_or_empty(
lambda: UserRepository(prisma_client).table.find_many(
where={"user_id": {"in": list(user_ids)}}, # mutable-ok: Prisma find_many where= is a dict
),
lambda: find_many_in(UserRepository(prisma_client).table, "user_id", user_ids),
"Failed user detail recovery for %d user ids: %s",
len(user_ids),
)

View file

@ -0,0 +1,145 @@
"""
Prisma `{"in": [...]}` filters whose value list may outgrow Postgres's bind-parameter cap.
A membership filter binds one parameter per value and Postgres caps a statement at 32,767,
so each operation here splits the deduplicated values into chunks of `chunk_size` values
(`IN_LIST_CHUNK_SIZE` by default, at most `MAX_IN_LIST_CHUNK_SIZE` so the rest of the filter
keeps headroom under the cap), runs them one after another (a transaction handle works as
`table`), and combines the results. An empty list returns without querying.
`not_in` cannot be chunked: a row must be outside every chunk at once. Such sites need
`<> ALL($1::text[])` in raw SQL or a relation filter instead.
"""
from collections.abc import Awaitable, Callable, Hashable, Iterable, Mapping
from itertools import accumulate, chain, repeat, takewhile
from typing import Final, Literal, TypeAlias, TypeVar
from litellm.repositories.prisma_protocols import CountTable, DeleteManyTable, FindManyTable, UpdateManyTable
IN_LIST_CHUNK_SIZE: Final = 5_000
MAX_IN_LIST_CHUNK_SIZE: Final = 30_000
LOGICAL_KEYS: Final = frozenset({"AND", "OR", "NOT"})
RowT: Final = TypeVar("RowT")
ResultT: Final = TypeVar("ResultT")
Atomicity: TypeAlias = Literal["caller_transaction", "per_chunk_ok"]
"""More than `chunk_size` values means more than one statement. `caller_transaction`
states `table` is a transaction handle, so the chunks commit together; `per_chunk_ok` states
the caller accepts earlier chunks staying applied when a later one fails."""
class SameFieldFilterError(ValueError):
pass
class ChunkedFieldWriteError(ValueError):
"""An update that writes the chunked field can move a row into a later chunk, which then updates it again."""
def _as_clauses(value: object) -> tuple[object, ...]:
match value:
case list() | tuple():
return tuple(value) # pyright: ignore[reportUnknownVariableType, reportUnknownArgumentType] # filters nest arbitrary data
case _:
return (value,)
def _logical_clauses(clause: object) -> tuple[object, ...]:
match clause:
case Mapping():
return tuple(chain.from_iterable(_as_clauses(clause[key]) for key in LOGICAL_KEYS if key in clause)) # pyright: ignore[reportUnknownArgumentType] # filters nest arbitrary data
case _:
return ()
def _filters_field(where: Mapping[str, object], field: str) -> bool:
"""Whether `field` is filtered in `where` or in any AND / OR / NOT clause under it, walked level by level."""
levels: Final = accumulate(
repeat(None),
lambda level, _: tuple(chain.from_iterable(map(_logical_clauses, level))),
initial=(where,),
)
return any(
isinstance(clause, Mapping) and field in clause for clause in chain.from_iterable(takewhile(bool, levels))
)
def _chunk_filter(field: str, chunk: tuple[Hashable, ...], where: Mapping[str, object] | None) -> Mapping[str, object]:
membership: Final = {field: {"in": list(chunk)}} # mutable-ok: the dict and list a hand-written filter sends
if where is None:
return membership
return {"AND": (dict(where), membership)} # mutable-ok: prisma's query builder only accepts dict filters
async def _each_chunk(
field: str,
values: Iterable[Hashable],
where: Mapping[str, object] | None,
run: Callable[[Mapping[str, object]], Awaitable[ResultT]],
chunk_size: int,
) -> tuple[ResultT, ...]:
if not 1 <= chunk_size <= MAX_IN_LIST_CHUNK_SIZE:
raise ValueError(f"chunk_size must be between 1 and {MAX_IN_LIST_CHUNK_SIZE:,}, got {chunk_size}")
if where is not None and _filters_field(where, field):
raise SameFieldFilterError(f"`where` already filters `{field}`; fold that condition into the values instead")
unique: Final = tuple(dict.fromkeys(values))
starts: Final = range(0, len(unique), chunk_size)
return tuple([await run(_chunk_filter(field, unique[start : start + chunk_size], where)) for start in starts])
async def find_many_in(
table: FindManyTable[RowT],
field: str,
values: Iterable[Hashable],
*,
where: Mapping[str, object] | None = None,
chunk_size: int = IN_LIST_CHUNK_SIZE,
) -> tuple[RowT, ...]:
"""Rows in chunk order. No take/skip/cursor/order/distinct: none of them survive a split."""
pages: Final = await _each_chunk(field, values, where, lambda chunk: table.find_many(where=chunk), chunk_size)
return tuple(chain.from_iterable(pages))
async def count_in(
table: CountTable,
field: str,
values: Iterable[Hashable],
*,
where: Mapping[str, object] | None = None,
chunk_size: int = IN_LIST_CHUNK_SIZE,
) -> int:
return sum(await _each_chunk(field, values, where, lambda chunk: table.count(where=chunk), chunk_size))
async def update_many_in(
table: UpdateManyTable,
field: str,
values: Iterable[Hashable],
*,
data: Mapping[str, object],
atomicity: Atomicity,
where: Mapping[str, object] | None = None,
chunk_size: int = IN_LIST_CHUNK_SIZE,
) -> int:
if field in data:
raise ChunkedFieldWriteError(
f"`data` writes `{field}`, the chunked field; a row it moves can match a later chunk"
)
payload: Final = dict(data) # mutable-ok: prisma's query builder only accepts dict payloads
return sum(
await _each_chunk(field, values, where, lambda chunk: table.update_many(data=payload, where=chunk), chunk_size)
)
async def delete_many_in(
table: DeleteManyTable,
field: str,
values: Iterable[Hashable],
*,
atomicity: Atomicity,
where: Mapping[str, object] | None = None,
chunk_size: int = IN_LIST_CHUNK_SIZE,
) -> int:
return sum(await _each_chunk(field, values, where, lambda chunk: table.delete_many(where=chunk), chunk_size))

View file

@ -118,6 +118,22 @@ class SpendLinkedTable(Protocol[RowT_co]):
async def update_many(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> int: ...
class FindManyTable(Protocol[RowT_co]):
async def find_many(self, *, where: Mapping[str, object]) -> Sequence[RowT_co]: ...
class CountTable(Protocol):
async def count(self, *, where: Mapping[str, object]) -> int: ...
class UpdateManyTable(Protocol):
async def update_many(self, *, data: Mapping[str, object], where: Mapping[str, object]) -> int: ...
class DeleteManyTable(Protocol):
async def delete_many(self, *, where: Mapping[str, object]) -> int: ...
class BatchTable(Protocol):
def update(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> None: ...

View file

@ -0,0 +1,502 @@
#!/usr/bin/env python3
"""Fail CI on SQL `IN (...)` lists whose length nothing bounds (Postgres caps a statement at 32,767 binds).
Reported under litellm/ and enterprise/: a Prisma `"in"` / `"not_in"` filter over a value with no
fixed size, and a raw-SQL `IN (` followed by a value spliced in at runtime. Chunk an `in` list with
`litellm.repositories.chunked_in`, pass raw SQL one array parameter, or record a real bound with
`# bounded-ok: <reason>` on the reported line or the line above.
Existing findings live in `unbounded_in_baseline.txt`, keyed without line numbers. A finding the
baseline lacks fails the run, as does an entry no finding matches; `--update-baseline` rewrites it.
Usage: python check_unbounded_in_lists.py [--update-baseline] [--baseline FILE] [files-or-dirs...]
"""
from __future__ import annotations
import argparse
import ast
import io
import re
import sys
import tokenize
from collections.abc import Callable, Iterable, Iterator, Mapping
from dataclasses import dataclass
from functools import reduce
from pathlib import Path
from typing import Final
REPO_ROOT: Final = Path(__file__).resolve().parents[2]
DEFAULT_TARGETS: Final = ("litellm", "enterprise")
DEFAULT_BASELINE: Final = Path(__file__).resolve().with_name("unbounded_in_baseline.txt")
EXEMPT_PATHS: Final = frozenset({"litellm/repositories/chunked_in.py"})
MODULE_SCOPE: Final = "<module>"
BASELINE_HEADER: Final = (
"# Grandfathered findings of check_unbounded_in_lists.py: path::scope::kind::subject::occurrence.\n"
"# Fix a site and delete its line; regenerate with `check_unbounded_in_lists.py --update-baseline`.\n"
)
MEMBERSHIP_KEYS: Final = frozenset({"in", "not_in", "notIn"})
TYPED_DICT_MODULES: Final = frozenset({"typing", "typing_extensions"})
CONSTANT_WRAPPERS: Final = frozenset({"list", "tuple", "sorted", "frozenset", "set"})
FREEZING_WRAPPERS: Final = frozenset({"tuple", "frozenset"})
MIN_REASON_LEN: Final = 3
MARKER: Final = re.compile(r"#\s*bounded-ok(?::[ \t]*(?P<reason>[^#]*))?")
# The text right after `IN (` is where a runtime value lands: an f-string or format
# slot (`{x}`, never the escaped `{{`), a `%` slot, or the end of the literal itself.
SPLICED_IN: Final = re.compile(r"\bIN\s*\(\s*(?:\{(?!\{)|%s\b|%\(|$)", re.IGNORECASE)
IN_OPERAND: Final = re.compile(r"(\S+)\s+(?:NOT\s+)?$", re.IGNORECASE)
STRING_PREFIX_AND_QUOTES: Final = re.compile(r"^[rbfuRBFU]{0,2}(?=[\"'])|[\\\"']")
CLOSING_QUOTES: Final = re.compile(r"(?:\"\"\"|'''|\"|')$")
@dataclass(frozen=True, slots=True)
class Finding:
path: Path
line: int
kind: str
message: str
scope: str = MODULE_SCOPE
subject: str = ""
value: str = ""
def render(self) -> str:
return f"{self.path}:{self.line}: {self.kind} {self.message}"
@dataclass(frozen=True, slots=True)
class Marker:
reason: str
standalone: bool
@property
def valid(self) -> bool:
return len(self.reason) >= MIN_REASON_LEN
@dataclass(frozen=True, slots=True)
class Markers:
by_line: Mapping[int, Marker]
def exempt(self, line: int) -> bool:
"""A marker on the line itself, or alone on the line above it, speaks for it."""
same: Final = self.by_line.get(line)
above: Final = self.by_line.get(line - 1)
return (same is not None and same.valid) or (above is not None and above.standalone and above.valid)
def read_markers(source: str) -> Markers:
try:
tokens: Final = tuple(tokenize.generate_tokens(io.StringIO(source).readline))
except (tokenize.TokenError, SyntaxError):
return Markers({})
return Markers(
{
token.start[0]: Marker(
reason=(match.group("reason") or "").strip(),
standalone=not token.line[: token.start[1]].strip(),
)
for token in tokens
if token.type == tokenize.COMMENT
for match in (MARKER.search(token.string),)
if match is not None
}
)
def _fixed_element(element: ast.expr, constants: frozenset[str]) -> bool:
match element:
case ast.Starred(value=value):
return has_fixed_size(value, constants)
case _:
return True
def has_fixed_size(value: ast.expr, constants: frozenset[str]) -> bool:
"""Whether the value's length is visible in the source rather than decided at runtime."""
match value:
case ast.List(elts=elts) | ast.Tuple(elts=elts) | ast.Set(elts=elts):
return all(_fixed_element(elt, constants) for elt in elts)
case ast.Constant():
return True
case ast.Name(id=name):
return name in constants
case ast.Call(func=ast.Name(id=wrapper), args=[argument], keywords=[]) if wrapper in CONSTANT_WRAPPERS:
return has_fixed_size(argument, constants)
case _:
return False
def _module_binding(stmt: ast.stmt) -> tuple[tuple[str, ast.expr], ...]:
match stmt:
case ast.Assign(targets=[ast.Name(id=name)], value=value):
return ((name, value),)
case ast.AnnAssign(target=ast.Name(id=name), value=ast.expr() as value):
return ((name, value),)
case _:
return ()
def _stays_fixed(value: ast.expr, constants: frozenset[str]) -> bool:
"""has_fixed_size, less the shapes a later append or extend could grow."""
match value:
case ast.Tuple(elts=elts):
return all(_fixed_element(elt, constants) for elt in elts)
case ast.Constant():
return True
case ast.Name(id=name):
return name in constants
case ast.Call(func=ast.Name(id=wrapper), args=[argument], keywords=[]) if wrapper in FREEZING_WRAPPERS:
return has_fixed_size(argument, constants)
case _:
return False
def module_constants(tree: ast.Module) -> frozenset[str]:
"""Module-level names bound exactly once to a frozen value of fixed size, in binding
order so one constant may be built from another. Casing plays no part: an ALL_CAPS
name that is imported or filled at runtime is as unbounded as any other."""
bound: Final = tuple(binding for stmt in tree.body for binding in _module_binding(stmt))
names: Final = tuple(name for name, _ in bound)
rebound: Final = frozenset(name for name in names if names.count(name) > 1)
def fold(constants: frozenset[str], binding: tuple[str, ast.expr]) -> frozenset[str]:
name, value = binding
return constants | {name} if name not in rebound and _stays_fixed(value, constants) else constants
return reduce(fold, bound, frozenset())
@dataclass(frozen=True, slots=True)
class Span:
start: int
end: int
qualname: str
def _spans(node: ast.AST, prefix: str) -> Iterator[Span]:
for child in ast.iter_child_nodes(node):
match child:
case ast.FunctionDef(name=name) | ast.AsyncFunctionDef(name=name) | ast.ClassDef(name=name):
yield Span(child.lineno, child.end_lineno or child.lineno, prefix + name)
yield from _spans(child, f"{prefix}{name}.")
case _:
yield from _spans(child, prefix)
def scope_finder(tree: ast.AST) -> Callable[[int], str]:
"""The innermost function or class around a line, dotted like a qualname, else `<module>`."""
spans: Final = tuple(_spans(tree, ""))
def scope_of(line: int) -> str:
enclosing: Final = tuple(span for span in spans if span.start <= line <= span.end)
return max(enclosing, key=lambda span: (span.start, -span.end)).qualname if enclosing else MODULE_SCOPE
return scope_of
def _field_name(key: ast.expr) -> str:
match key:
case ast.Constant(value=str(name)):
return name
case _:
return f"[{ast.unparse(key)}]"
def _field_bindings(node: ast.AST) -> Iterator[tuple[str, ast.expr]]:
"""Where a dict literal is written as a field's filter: `{field: {...}}`, `where[field] = {...}`
or `Filter(field={...})`. A computed field reads as `[expr]`."""
match node:
case ast.Dict(keys=keys, values=values):
yield from ((_field_name(key), value) for key, value in zip(keys, values) if key is not None)
case ast.Assign(targets=[ast.Subscript(slice=key)], value=value):
yield (_field_name(key), value)
case ast.Call(keywords=keywords):
yield from ((keyword.arg, keyword.value) for keyword in keywords if keyword.arg is not None)
case _:
return
def _filtered_fields(tree: ast.AST) -> Mapping[int, str]:
"""id() of each dict literal written as a field's filter, mapped to that field."""
return {
id(value): field
for node in ast.walk(tree)
for field, value in _field_bindings(node)
if isinstance(value, ast.Dict)
}
def _is_typed_dict(func: ast.expr) -> bool:
match func:
case ast.Name(id="TypedDict"):
return True
case ast.Attribute(value=ast.Name(id=module), attr="TypedDict"):
return module in TYPED_DICT_MODULES
case _:
return False
def _typed_dict_field_map(node: ast.AST) -> ast.expr | None:
"""The field map of a functional `TypedDict("Name", {...})`, whose keys are field names, not filters."""
match node:
case ast.Call(func=func, args=[_, fields, *_]) if _is_typed_dict(func):
return fields
case ast.Call(func=func, keywords=keywords) if _is_typed_dict(func):
return next((keyword.value for keyword in keywords if keyword.arg == "fields"), None)
case _:
return None
def _typed_dict_field_maps(tree: ast.AST) -> frozenset[int]:
"""id() of each dict literal passed as a functional TypedDict's field map."""
return frozenset(id(fields) for fields in map(_typed_dict_field_map, ast.walk(tree)) if fields is not None)
def _prisma_advice(key: str) -> str:
if key == "in":
return (
"Chunk it with `litellm.repositories.chunked_in` (find_many_in / count_in / update_many_in / "
"delete_many_in)"
)
return "A negated list cannot be chunked: use `<> ALL($1::text[])` in raw SQL or a relation filter"
def prisma_findings(path: Path, tree: ast.Module) -> Iterator[Finding]:
constants: Final = module_constants(tree)
scope_of: Final = scope_finder(tree)
fields: Final = _filtered_fields(tree)
typed_dict_field_maps: Final = _typed_dict_field_maps(tree)
for node in ast.walk(tree):
if not isinstance(node, ast.Dict) or id(node) in typed_dict_field_maps:
continue
for key, value in zip(node.keys, node.values):
if not (isinstance(key, ast.Constant) and key.value in MEMBERSHIP_KEYS):
continue
if has_fixed_size(value, constants):
continue
yield Finding(
path,
key.lineno,
"prisma",
f'`"{key.value}"` filter over `{ast.unparse(value)}` has no written bound: it binds one '
f"parameter per value and Postgres caps a statement at 32,767. {_prisma_advice(key.value)}, "
f"or record the bound with `# bounded-ok: <reason>`",
scope=scope_of(key.lineno),
subject=f"{fields.get(id(node), '?')}.{key.value}",
value=_normalized(ast.unparse(value)),
)
def _literal_body(lines: tuple[bytes, ...], node: ast.expr) -> str | None:
"""The literal's source text with its closing quotes removed, so a literal that
ends right after `IN (` reads as an open list rather than as `IN ('`. Column
offsets count UTF-8 bytes, so the slice is taken on the encoded lines."""
end_line: Final = node.end_lineno
end_col: Final = node.end_col_offset
if end_line is None or end_col is None:
return None
first: Final = node.lineno - 1
last: Final = end_line - 1
segment: Final = (
lines[first][node.col_offset : end_col]
if first == last
else b"".join((lines[first][node.col_offset :], *lines[first + 1 : last], lines[last][:end_col]))
)
return CLOSING_QUOTES.sub("", segment.decode("utf-8", errors="replace"))
def _fstring_part_ids(tree: ast.AST) -> frozenset[int]:
"""ids() of the literal pieces inside f-strings, which the enclosing JoinedStr already covers."""
return frozenset(
id(part)
for node in ast.walk(tree)
if isinstance(node, ast.JoinedStr)
for value in node.values
for part in (
(value,)
if isinstance(value, ast.Constant)
else tuple(ast.walk(value.format_spec))
if isinstance(value, ast.FormattedValue) and value.format_spec is not None
else ()
)
)
def _normalized(text: str) -> str:
return " ".join(text.split())
def _slot_end(body: str, start: int) -> int:
"""Just past the `)` closing an `IN (` slot, or the end of the literal when it has none."""
close: Final = body.find(")", start)
return len(body) if close == -1 else close + 1
def raw_sql_findings(path: Path, source: str, tree: ast.AST) -> Iterator[Finding]:
parts: Final = _fstring_part_ids(tree)
scope_of: Final = scope_finder(tree)
lines: Final = tuple(source.encode("utf-8").splitlines(keepends=True))
for node in ast.walk(tree):
is_text = isinstance(node, ast.JoinedStr) or (isinstance(node, ast.Constant) and isinstance(node.value, str))
if not is_text or id(node) in parts:
continue
body = _literal_body(lines, node)
match = None if body is None else SPLICED_IN.search(body)
if body is None or match is None:
continue
in_line = node.lineno + body[: match.start()].count("\n")
where = "" if in_line == node.lineno else f" (the `IN (` is on line {in_line})"
operand = IN_OPERAND.search(body[: match.start()])
yield Finding(
path,
node.lineno,
"raw-sql",
f"`IN (` takes a list spliced in at runtime{where}: it binds one parameter per value and Postgres "
f"caps a statement at 32,767. Pass the list as one array parameter (`= ANY($1::text[])`, or "
f"`<> ALL($1::text[])` for `NOT IN`), or record the bound with `# bounded-ok: <reason>`",
scope=scope_of(node.lineno),
subject=f"{STRING_PREFIX_AND_QUOTES.sub('', operand.group(1)) if operand else '?'}.IN",
value=_normalized(body[match.start() : _slot_end(body, match.end())]),
)
def marker_findings(path: Path, markers: Markers, scope_of: Callable[[int], str]) -> Iterator[Finding]:
for line, marker in sorted(markers.by_line.items()):
if not marker.valid:
yield Finding(
path,
line,
"marker",
"`# bounded-ok` needs a reason naming the bound: `# bounded-ok: <reason>`",
scope=scope_of(line),
subject="bounded-ok",
)
def check_file(path: Path) -> tuple[Finding, ...]:
try:
source: Final = path.read_text(encoding="utf-8")
tree: Final = ast.parse(source, filename=str(path))
except (OSError, UnicodeDecodeError, SyntaxError) as exc:
return (Finding(path, getattr(exc, "lineno", None) or 0, "unreadable", str(exc)),)
markers: Final = read_markers(source)
return (
*marker_findings(path, markers, scope_finder(tree)),
*(
finding
for finding in (*prisma_findings(path, tree), *raw_sql_findings(path, source, tree))
if not markers.exempt(finding.line)
),
)
def collect_paths(raw: Iterable[str]) -> Iterator[Path]:
for item in raw:
path = Path(item)
if path.is_dir():
yield from sorted(path.rglob("*.py"))
elif path.suffix == ".py":
yield path
def repo_relative(path: Path) -> str:
resolved: Final = path.resolve()
return resolved.relative_to(REPO_ROOT).as_posix() if resolved.is_relative_to(REPO_ROOT) else resolved.as_posix()
def scan(paths: Iterable[Path]) -> tuple[Finding, ...]:
return tuple(
sorted(
(f for path in paths if repo_relative(path) not in EXEMPT_PATHS for f in check_file(path)),
key=lambda f: (str(f.path), f.line, f.kind),
)
)
def identify(findings: tuple[Finding, ...]) -> Mapping[str, Finding]:
"""Each finding keyed by `path scope kind subject `value` occurrence`, the value being the
filtered expression's source and the occurrence counting the earlier findings in the same file
that share the rest of the key. No line number goes in, so code shifting up or down leaves the
key alone, while a different expression on the same field reads as a new finding."""
ordered: Final = sorted(findings, key=lambda f: (str(f.path), f.line))
keys: Final = tuple(
f"{repo_relative(f.path)} {f.scope} {f.kind} {f.subject or '-'}" + (f" `{f.value}`" if f.value else "")
for f in ordered
)
return {f"{key} {keys[:index].count(key)}": finding for index, (key, finding) in enumerate(zip(keys, ordered))}
def read_baseline(path: Path) -> frozenset[str]:
if not path.exists():
return frozenset()
return frozenset(
stripped
for line in path.read_text(encoding="utf-8").splitlines()
for stripped in (line.strip(),)
if stripped and not stripped.startswith("#")
)
def covered_by(targets: tuple[str, ...]) -> Callable[[str], bool]:
"""Whether a baseline entry's file lies under one of the scanned targets."""
roots: Final = tuple(repo_relative(Path(target)) for target in targets)
def covers(entry: str) -> bool:
entry_path: Final = entry.split(" ", 1)[0]
return any(entry_path == root or entry_path.startswith(f"{root}/") for root in roots)
return covers
@dataclass(frozen=True, slots=True)
class Options:
targets: tuple[str, ...]
baseline: Path
update_baseline: bool
def parse_options(argv: Iterable[str]) -> Options:
parser: Final = argparse.ArgumentParser(description="Fail on SQL IN lists with no written bound.")
parser.add_argument("targets", nargs="*", default=list(DEFAULT_TARGETS))
parser.add_argument("--baseline", default=str(DEFAULT_BASELINE))
parser.add_argument("--update-baseline", action="store_true")
namespace: Final = parser.parse_args(list(argv))
return Options(
targets=tuple(str(target) for target in namespace.targets),
baseline=Path(str(namespace.baseline)),
update_baseline=bool(namespace.update_baseline),
)
def main(argv: Iterable[str]) -> int:
options: Final = parse_options(argv)
findings: Final = scan(collect_paths(options.targets))
current: Final = identify(findings)
baseline: Final = read_baseline(options.baseline)
covers: Final = covered_by(options.targets)
if options.update_baseline:
entries: Final = sorted({*(entry for entry in baseline if not covers(entry)), *current})
options.baseline.write_text(BASELINE_HEADER + "".join(f"{entry}\n" for entry in entries), encoding="utf-8")
print(f"Wrote {len(entries)} baseline entries to {options.baseline}")
return 0
new: Final = tuple(finding for key, finding in current.items() if key not in baseline)
stale: Final = sorted(entry for entry in baseline if covers(entry) and entry not in current)
for finding in new:
print(finding.render())
for entry in stale:
print(f"{options.baseline}: stale entry `{entry}`: no finding matches it any more, delete the line")
counts: Final = {
kind: sum(1 for f in findings if f.kind == kind) for kind in ("prisma", "raw-sql", "marker", "unreadable")
}
summary: Final = ", ".join(f"{count} {kind}" for kind, count in counts.items() if count)
print(
f"\n{len(findings)} unbounded IN list(s) ({summary or 'none'}): {len(findings) - len(new)} baselined, "
f"{len(new)} new, {len(stale)} stale baseline entries."
)
return 1 if new or stale else 0
if __name__ == "__main__":
raise SystemExit(main(sys.argv[1:]))

View file

@ -0,0 +1,158 @@
# Grandfathered findings of check_unbounded_in_lists.py: path::scope::kind::subject::occurrence.
# Fix a site and delete its line; regenerate with `check_unbounded_in_lists.py --update-baseline`.
enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py CheckResponsesCost.check_responses_cost prisma id.in `[job.id for job in completed_jobs]` 0
enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py list_projects prisma team_id.in `user_team_ids` 0
litellm/integrations/shadow_eval_logger.py ShadowEvalLogger._active_jobs prisma job_id.in `[str(record.id) for record in records]` 0
litellm/llms/litellm_proxy/skills/handler.py LiteLLMSkillsHandler.list_skills prisma created_by.in `owner_scopes` 0
litellm/proxy/_experimental/mcp_server/db.py get_mcp_servers prisma server_id.in `server_ids` 0
litellm/proxy/_experimental/mcp_server/db.py get_user_env_vars_bulk prisma server_id.in `ids` 0
litellm/proxy/_experimental/mcp_server/db.py purge_user_oauth_credentials_for_server prisma user_id.in `[row.user_id for row in oauth_rows]` 0
litellm/proxy/_experimental/mcp_server/oauth2_flow_backfill.py backfill_null_oauth2_flows prisma server_id.in `server_ids_for_flow` 0
litellm/proxy/_experimental/mcp_server/oauth2_flow_backfill.py backfill_null_oauth2_flows prisma server_id.in `server_ids` 0
litellm/proxy/_experimental/mcp_server/toolset_db.py list_mcp_toolsets prisma toolset_id.in `toolset_ids` 0
litellm/proxy/agent_endpoints/endpoints.py _attach_keys_to_agents prisma agent_id.in `agent_ids` 0
litellm/proxy/agent_endpoints/endpoints.py get_agent_daily_activity prisma agent_id.in `list(agent_ids_list)` 0
litellm/proxy/agent_endpoints/endpoints.py get_agents prisma agent_id.in `agent_ids` 0
litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_skill_access.py SkillVisibility.where prisma name.in `sorted(self.granted)` 0
litellm/proxy/auth/auth_checks.py _fetch_uncached_model_access_group_budgets prisma access_group_name.in `list(uncached_groups)` 0
litellm/proxy/auth/auth_checks.py _fetch_uncached_tags prisma tag_name.in `list(tags_to_fetch)` 0
litellm/proxy/auth/auth_checks.py get_jwt_key_mapping_cache_keys_for_tokens prisma token.in `tuple(hashed_tokens)` 0
litellm/proxy/auth/auth_checks.py get_managed_vector_store_rows_by_uuids prisma vector_store_id.in `cache_misses` 0
litellm/proxy/common_utils/reset_budget_job.py _budget_link_where prisma budget_id.in `list(budget_ids)` 0
litellm/proxy/container_endpoints/ownership.py _get_allowed_container_ids prisma created_by.in `owner_scopes` 0
litellm/proxy/db/tool_registry_writer.py get_tools_by_names prisma tool_name.in `tool_names` 0
litellm/proxy/guardrails/guardrail_endpoints.py list_guardrail_submissions prisma team_id.in `visible_team_ids` 0
litellm/proxy/guardrails/usage_endpoints.py _build_usage_logs_where prisma ?.in `guardrail_ids` 0
litellm/proxy/guardrails/usage_endpoints.py guardrails_usage_detail prisma guardrail_id.in `metric_ids` 0
litellm/proxy/guardrails/usage_endpoints.py guardrails_usage_detail prisma guardrail_id.in `metric_ids` 1
litellm/proxy/guardrails/usage_endpoints.py guardrails_usage_detail prisma guardrail_id.in `metric_ids` 2
litellm/proxy/guardrails/usage_endpoints.py guardrails_usage_logs prisma request_id.in `request_ids` 0
litellm/proxy/list_api/list_framework.py _render raw-sql {field}.IN `IN ({placeholders})` 0
litellm/proxy/management_endpoints/access_group_endpoints.py _require_teams_exist prisma team_id.in `team_ids` 0
litellm/proxy/management_endpoints/access_group_endpoints.py _teams_touching prisma team_id.in `stored_team_ids` 0
litellm/proxy/management_endpoints/auto_router_endpoints.py _with_target_labels prisma team_id.in `list(team_ids)` 0
litellm/proxy/management_endpoints/auto_router_endpoints.py _with_target_labels prisma token.in `list(tokens)` 0
litellm/proxy/management_endpoints/auto_router_endpoints.py _with_target_labels prisma user_id.in `list(user_ids)` 0
litellm/proxy/management_endpoints/auto_router_endpoints.py get_shadow_eval_job prisma job_id.in `leg_ids` 0
litellm/proxy/management_endpoints/auto_router_endpoints.py start_shadow_eval prisma target_id.in `list(ids)` 0
litellm/proxy/management_endpoints/auto_router_endpoints.py start_shadow_eval prisma team_id.in `list(data.team_ids)` 0
litellm/proxy/management_endpoints/auto_router_endpoints.py start_shadow_eval prisma token.in `list(data.api_key_ids)` 0
litellm/proxy/management_endpoints/auto_router_endpoints.py start_shadow_eval prisma user_id.in `list(data.user_ids)` 0
litellm/proxy/management_endpoints/budget_management_endpoints.py info_budget prisma budget_id.in `data.budgets` 0
litellm/proxy/management_endpoints/common_daily_activity.py _build_aggregated_where_clause raw-sql api_key.IN `IN ({placeholders})` 0
litellm/proxy/management_endpoints/common_daily_activity.py _build_aggregated_where_clause raw-sql {entity_id_field}.IN `IN ({placeholders})` 0
litellm/proxy/management_endpoints/common_daily_activity.py _build_aggregated_where_clause raw-sql {entity_id_field}.IN `IN ({placeholders})` 1
litellm/proxy/management_endpoints/common_daily_activity.py _build_where_conditions prisma [entity_id_field].in `entity_id` 0
litellm/proxy/management_endpoints/common_daily_activity.py _build_where_conditions prisma api_key.in `api_key` 0
litellm/proxy/management_endpoints/common_daily_activity.py _build_where_conditions prisma not.in `exclude_entity_ids` 0
litellm/proxy/management_endpoints/common_daily_activity.py get_api_key_metadata prisma token.in `list(api_keys)` 0
litellm/proxy/management_endpoints/common_daily_activity.py get_api_key_metadata prisma token.in `list(missing_keys)` 0
litellm/proxy/management_endpoints/common_utils.py _team_admin_can_invite_user prisma team_id.in `admin_user_obj.teams` 0
litellm/proxy/management_endpoints/common_utils.py _user_has_admin_privileges prisma team_id.in `user_obj.teams` 0
litellm/proxy/management_endpoints/customer_endpoints.py delete_end_user prisma user_id.in `data.user_ids` 0
litellm/proxy/management_endpoints/customer_endpoints.py delete_end_user prisma user_id.in `data.user_ids` 1
litellm/proxy/management_endpoints/customer_endpoints.py get_customer_daily_activity prisma user_id.in `list(end_user_ids_list)` 0
litellm/proxy/management_endpoints/internal_user_endpoints.py _check_user_info_v2_access prisma team_id.in `caller_user.teams` 0
litellm/proxy/management_endpoints/internal_user_endpoints.py _resolve_user_email_metadata prisma user_id.in `list(user_ids)` 0
litellm/proxy/management_endpoints/internal_user_endpoints.py delete_user prisma created_by.in `data.user_ids` 0
litellm/proxy/management_endpoints/internal_user_endpoints.py delete_user prisma team_id.in `user_row.teams` 0
litellm/proxy/management_endpoints/internal_user_endpoints.py delete_user prisma updated_by.in `data.user_ids` 0
litellm/proxy/management_endpoints/internal_user_endpoints.py delete_user prisma user_id.in `data.user_ids` 0
litellm/proxy/management_endpoints/internal_user_endpoints.py delete_user prisma user_id.in `data.user_ids` 1
litellm/proxy/management_endpoints/internal_user_endpoints.py delete_user prisma user_id.in `data.user_ids` 2
litellm/proxy/management_endpoints/internal_user_endpoints.py delete_user prisma user_id.in `data.user_ids` 3
litellm/proxy/management_endpoints/internal_user_endpoints.py delete_user prisma user_id.in `data.user_ids` 4
litellm/proxy/management_endpoints/internal_user_endpoints.py delete_user prisma user_id.in `data.user_ids` 5
litellm/proxy/management_endpoints/internal_user_endpoints.py get_users prisma organization_id.in `org_id_list` 0
litellm/proxy/management_endpoints/internal_user_endpoints.py get_users prisma sso_user_id.in `sso_id_list` 0
litellm/proxy/management_endpoints/internal_user_endpoints.py get_users prisma user_id.in `user_id_list` 0
litellm/proxy/management_endpoints/internal_user_endpoints.py ui_view_users prisma organization_id.in `org_filter_ids` 0
litellm/proxy/management_endpoints/key_management_endpoints.py _apply_non_admin_alias_scope raw-sql team_id.IN `IN ({team_placeholders})` 0
litellm/proxy/management_endpoints/key_management_endpoints.py _build_key_filter_conditions prisma team_id.in `admin_team_ids` 0
litellm/proxy/management_endpoints/key_management_endpoints.py _build_key_filter_conditions prisma team_id.in `member_only_team_ids` 0
litellm/proxy/management_endpoints/key_management_endpoints.py _build_key_filter_conditions prisma team_id.in `member_team_ids` 0
litellm/proxy/management_endpoints/key_management_endpoints.py _fetch_user_team_objects prisma team_id.in `complete_user_info.teams` 0
litellm/proxy/management_endpoints/key_management_endpoints.py _list_key_helper prisma user_id.in `all_ids` 0
litellm/proxy/management_endpoints/key_management_endpoints.py bulk_update_team_keys prisma token.in `hashed_key_ids` 0
litellm/proxy/management_endpoints/key_management_endpoints.py delete_key_aliases prisma key_alias.in `key_aliases` 0
litellm/proxy/management_endpoints/key_management_endpoints.py delete_verification_tokens prisma token.in `hashed_tokens` 0
litellm/proxy/management_endpoints/key_management_endpoints.py info_key_fn_v2 prisma key_alias.in `data.key_aliases` 0
litellm/proxy/management_endpoints/mcp_management_endpoints.py fetch_all_mcp_servers prisma server_id.in `byok_server_ids` 0
litellm/proxy/management_endpoints/model_access_group_management_endpoints.py update_deployments_with_access_group prisma model_name.in `model_names` 0
litellm/proxy/management_endpoints/model_management_endpoints.py delete_team_models prisma model_id.in `model_ids` 0
litellm/proxy/management_endpoints/organization_endpoints.py deprecated_info_organization prisma organization_id.in `data.organizations` 0
litellm/proxy/management_endpoints/organization_endpoints.py get_organization_daily_activity prisma organization_id.in `list(org_ids_list)` 0
litellm/proxy/management_endpoints/organization_endpoints.py list_organization prisma organization_id.in `membership_org_ids` 0
litellm/proxy/management_endpoints/router_weights.py validate_router_settings_weights prisma model_id.in `list(deployment_ids)` 0
litellm/proxy/management_endpoints/session_endpoints.py revoke_ui_session_keys prisma token.in `revoked_tokens` 0
litellm/proxy/management_endpoints/tag_management_endpoints.py _get_model_names prisma model_id.in `model_ids` 0
litellm/proxy/management_endpoints/tag_management_endpoints.py _get_tag_list_scope prisma api_key.in `scoped_api_keys` 0
litellm/proxy/management_endpoints/tag_management_endpoints.py info_tag prisma tag_name.in `data.names` 0
litellm/proxy/management_endpoints/tag_management_endpoints.py list_tags prisma tag_name.in `used_tag_names` 0
litellm/proxy/management_endpoints/team_endpoints.py _append_permissions_to_specific_teams prisma team_id.in `team_ids` 0
litellm/proxy/management_endpoints/team_endpoints.py _authorize_and_filter_teams prisma organization_id.in `allowed_org_ids` 0
litellm/proxy/management_endpoints/team_endpoints.py _batch_resolve_access_group_resources prisma access_group_id.in `unique_ids` 0
litellm/proxy/management_endpoints/team_endpoints.py _build_team_list_where_conditions prisma organization_id.in `org_admin_org_ids` 0
litellm/proxy/management_endpoints/team_endpoints.py _build_team_list_where_conditions prisma organization_id.in `org_admin_org_ids` 1
litellm/proxy/management_endpoints/team_endpoints.py _build_team_list_where_conditions prisma team_id.in `list(own_team_ids)` 0
litellm/proxy/management_endpoints/team_endpoints.py _build_team_list_where_conditions prisma team_id.in `user_team_ids` 0
litellm/proxy/management_endpoints/team_endpoints.py _get_keys_count_by_team prisma team_id.in `page_team_ids` 0
litellm/proxy/management_endpoints/team_endpoints.py _hydrate_member_user_details prisma user_id.in `sorted(user_ids)` 0
litellm/proxy/management_endpoints/team_endpoints.py _resolve_existing_member_user_ids prisma user_id.in `sorted(requested_user_ids)` 0
litellm/proxy/management_endpoints/team_endpoints.py _resolve_team_daily_activity_scope prisma team_id.in `list(team_ids_list)` 0
litellm/proxy/management_endpoints/team_endpoints.py _sweep_deleted_team_references prisma team_id.in `tuple(team_ids)` 0
litellm/proxy/management_endpoints/team_endpoints.py _sweep_deleted_team_references_tx prisma team_id.in `tuple(team_ids)` 0
litellm/proxy/management_endpoints/team_endpoints.py _team_key_search_where prisma ?.in `tuple(scope.team_ids)` 0
litellm/proxy/management_endpoints/team_endpoints.py _team_key_search_where prisma ?.in `tuple(scope.team_ids)` 1
litellm/proxy/management_endpoints/team_endpoints.py _team_key_search_where prisma ?.notIn `tuple(scope.exclude_team_ids)` 0
litellm/proxy/management_endpoints/team_endpoints.py _team_key_search_where prisma ?.notIn `tuple(scope.exclude_team_ids)` 1
litellm/proxy/management_endpoints/team_endpoints.py _team_key_search_where prisma token.in `own_keys` 0
litellm/proxy/management_endpoints/team_endpoints.py _team_key_search_where prisma token.in `own_keys` 1
litellm/proxy/management_endpoints/team_endpoints.py _team_member_delete prisma user_id.in `sorted(addressed_user_ids)` 0
litellm/proxy/management_endpoints/team_endpoints.py _team_member_delete prisma user_id.in `sorted(user_ids_to_delete)` 0
litellm/proxy/management_endpoints/team_endpoints.py _team_member_delete prisma user_id.in `sorted(user_ids_to_delete)` 1
litellm/proxy/management_endpoints/team_endpoints.py _team_user_spend_sql raw-sql sl.team_id.IN `IN ({team_placeholders})` 0
litellm/proxy/management_endpoints/team_endpoints.py delete_team prisma team_id.in `data.team_ids` 0
litellm/proxy/management_endpoints/team_endpoints.py delete_team prisma team_id.in `data.team_ids` 1
litellm/proxy/management_endpoints/team_endpoints.py get_all_team_memberships prisma team_id.in `team_ids` 0
litellm/proxy/management_endpoints/team_endpoints.py list_available_teams prisma team_id.in `available_teams` 0
litellm/proxy/management_endpoints/tool_management_endpoints.py get_tool_spend prisma tool_name.in `[row.tool_name for row in top_tools]` 0
litellm/proxy/management_endpoints/tool_management_endpoints.py get_tool_usage_logs prisma request_id.in `request_ids` 0
litellm/proxy/management_endpoints/ui_sso.py fetch_cli_sso_team_details prisma team_id.in `teams` 0
litellm/proxy/management_endpoints/user_agent_analytics_endpoints.py get_per_user_analytics prisma tag.in `tag_filters` 0
litellm/proxy/management_endpoints/user_agent_analytics_endpoints.py get_per_user_analytics prisma token.in `list(api_keys)` 0
litellm/proxy/management_endpoints/user_agent_analytics_endpoints.py get_per_user_analytics prisma user_id.in `user_ids` 0
litellm/proxy/management_endpoints/workflow_management_endpoints.py list_workflow_runs prisma ?.in `statuses` 0
litellm/proxy/management_helpers/bulk_user_creation.py _existing_user_conflicts prisma user_email.in `emails` 0
litellm/proxy/management_helpers/bulk_user_creation.py _existing_user_conflicts prisma user_id.in `user_ids` 0
litellm/proxy/management_helpers/bulk_user_creation.py _insert_users prisma user_id.in `list(requested)` 0
litellm/proxy/management_helpers/bulk_user_creation.py _load_teams prisma team_id.in `sorted(team_ids)` 0
litellm/proxy/management_helpers/bulk_user_creation.py _write_audit_logs prisma user_id.in `created_ids` 0
litellm/proxy/management_helpers/bulk_user_deletion.py _in_filter prisma [field].in `sorted(values)` 0
litellm/proxy/management_helpers/object_permission_utils.py _get_db_mcp_servers_by_identifiers prisma alias.in `identifier_list` 0
litellm/proxy/management_helpers/object_permission_utils.py _get_db_mcp_servers_by_identifiers prisma server_id.in `identifier_list` 0
litellm/proxy/management_helpers/object_permission_utils.py _get_db_mcp_servers_by_identifiers prisma server_name.in `identifier_list` 0
litellm/proxy/management_helpers/resource_display_names.py agent_display_names prisma agent_id.in `tuple(wanted)` 0
litellm/proxy/management_helpers/resource_display_names.py key_display_names prisma token.in `tuple(frozenset(tokens))` 0
litellm/proxy/management_helpers/resource_display_names.py mcp_server_display_names prisma server_id.in `tuple(wanted)` 0
litellm/proxy/policy_engine/policy_resolve_endpoints.py _build_alias_where prisma [field].in `exact` 0
litellm/proxy/policy_engine/policy_resolve_endpoints.py _find_affected_by_team_patterns prisma team_id.in `matched_team_ids` 0
litellm/proxy/proxy_server.py _add_access_group_models_to_team_models prisma access_group_id.in `list(all_access_group_ids)` 0
litellm/proxy/proxy_server.py _fetch_db_models_for_search prisma not.in `list(db_model_ids_in_router)` 0
litellm/proxy/proxy_server.py _gather_team_accessible_model_ids prisma model_name.in `_resolved_names` 0
litellm/proxy/proxy_server.py get_all_team_models prisma team_id.in `user_teams` 0
litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py _prune_filter prisma model.in `chunk` 0
litellm/proxy/spend_tracking/spend_management_endpoints.py _find_team_rows prisma team_id.in `team_ids` 0
litellm/proxy/spend_tracking/spend_management_endpoints.py ui_view_session_spend_logs prisma team_id.in `permitted_team_ids` 0
litellm/proxy/spend_tracking/spend_management_endpoints.py ui_view_spend_logs prisma team_id.in `permitted_team_ids` 0
litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py _validate_default_teams_exist prisma team_id.in `list(team_ids)` 0
litellm/proxy/utils.py PrismaClient.check_view_exists raw-sql viewname.IN `IN ( {expected_views_str} )` 0
litellm/proxy/utils.py PrismaClient.delete_data prisma team_id.in `team_id_list` 0
litellm/proxy/utils.py PrismaClient.delete_data prisma team_id.in `team_id_list` 1
litellm/proxy/utils.py PrismaClient.delete_data prisma token.in `hashed_tokens` 0
litellm/proxy/utils.py PrismaClient.delete_data prisma token.in `hashed_tokens` 1
litellm/proxy/utils.py PrismaClient.get_data prisma budget_id.in `budget_id_list` 0
litellm/proxy/utils.py PrismaClient.get_data prisma team_id.in `team_id_list` 0
litellm/proxy/utils.py PrismaClient.get_data prisma user_id.in `user_id_list` 0
litellm/proxy/utils.py prefetch_config_params prisma param_name.in `param_names` 0
litellm/router_utils/auto_router_model_naming.py <module> raw-sql classifier_type.IN `IN ({_LLM_CLASSIFIER_TYPES_SQL})` 0

View file

@ -0,0 +1,162 @@
import os
import uuid
from datetime import timedelta
from collections.abc import AsyncIterator
from contextlib import asynccontextmanager
from types import SimpleNamespace
from typing import Final
from urllib.parse import parse_qsl, urlencode, urlsplit, urlunsplit
import psycopg
import pytest
from prisma import Prisma
from prisma.errors import DataError
from psycopg import sql
from litellm.proxy.spend_tracking.key_metadata_recovery import attach_user_details
from litellm.repositories.chunked_in import count_in, delete_many_in, find_many_in, update_many_in
ROWS: Final = 40_000
OUTSIDE: Final = 25
def _scoped_url(url: str, schema: str) -> str:
parsed: Final = urlsplit(url)
return urlunsplit(parsed._replace(query=urlencode({**dict(parse_qsl(parsed.query)), "schema": schema})))
@asynccontextmanager
async def _user_table(users: int) -> AsyncIterator[Prisma]:
"""A private schema holding a copy of the migrated `LiteLLM_UserTable`, seeded with `users` rows."""
schema: Final = f"integration_{uuid.uuid4().hex}"
url: Final = os.environ["DATABASE_URL"]
table: Final = sql.Identifier(schema, "LiteLLM_UserTable")
with psycopg.connect(url, autocommit=True) as setup:
setup.execute(sql.SQL("CREATE SCHEMA {}").format(sql.Identifier(schema)))
try:
setup.execute(
sql.SQL('CREATE TABLE {} (LIKE "LiteLLM_UserTable" INCLUDING DEFAULTS INCLUDING CONSTRAINTS)').format(
table
)
)
setup.execute(
sql.SQL(
"INSERT INTO {} (user_id, user_email) "
"SELECT 'user-' || n, 'user-' || n || '@example.com' FROM generate_series(0, %s) n"
).format(table),
(users - 1,),
)
database: Final = Prisma(datasource={"url": _scoped_url(url, schema)})
await database.connect()
try:
yield database
finally:
await database.disconnect()
finally:
setup.execute(sql.SQL("DROP SCHEMA {} CASCADE").format(sql.Identifier(schema)))
@asynccontextmanager
async def _config_table() -> AsyncIterator[tuple[Prisma, str]]:
"""A private schema holding only `LiteLLM_Config`, seeded with ROWS listed and OUTSIDE unlisted rows."""
schema: Final = f"integration_{uuid.uuid4().hex}"
url: Final = os.environ["DATABASE_URL"]
scoped_url: Final = _scoped_url(url, schema)
table: Final = sql.Identifier(schema, "LiteLLM_Config")
with psycopg.connect(url, autocommit=True) as setup:
setup.execute(sql.SQL("CREATE SCHEMA {}").format(sql.Identifier(schema)))
try:
setup.execute(
sql.SQL(
"CREATE TABLE {} (param_name text PRIMARY KEY, param_value jsonb, "
"last_run_at timestamp(3), reload_revision bigint NOT NULL DEFAULT 0)"
).format(table)
)
setup.execute(
sql.SQL(
"INSERT INTO {} (param_name) SELECT 'listed-' || n FROM generate_series(0, %s) n "
"UNION ALL SELECT 'outside-' || n FROM generate_series(0, %s) n"
).format(table),
(ROWS - 1, OUTSIDE - 1),
)
database: Final = Prisma(datasource={"url": scoped_url})
await database.connect()
try:
yield database, schema
finally:
await database.disconnect()
finally:
setup.execute(sql.SQL("DROP SCHEMA {} CASCADE").format(sql.Identifier(schema)))
def _listed() -> list[str]:
return [f"listed-{n}" for n in range(ROWS)]
def _count(schema: str, condition: sql.Composable) -> int:
with psycopg.connect(os.environ["DATABASE_URL"]) as connection:
row: Final = connection.execute(
sql.SQL("SELECT count(*) FROM {} WHERE ").format(sql.Identifier(schema, "LiteLLM_Config")) + condition
).fetchone()
assert row is not None
return int(row[0])
@pytest.mark.covers("other.database.chunked_in.raw_in_list_over_bind_cap_fails")
async def test_a_raw_in_filter_over_the_bind_parameter_cap_is_rejected_by_postgres() -> None:
async with _config_table() as (database, schema):
where: Final = {"param_name": {"in": _listed()}}
with pytest.raises(DataError, match="too many bind variables"):
await database.litellm_config.count(where=where)
with pytest.raises(DataError, match="too many bind variables"):
await database.litellm_config.update_many(where=where, data={"reload_revision": 1})
with pytest.raises(DataError, match="too many bind variables"):
await database.litellm_config.delete_many(where=where)
assert _count(schema, sql.SQL("reload_revision = 0")) == ROWS + OUTSIDE
@pytest.mark.covers(
"other.database.chunked_in.find_many_in_returns_every_row",
"other.database.chunked_in.count_in_counts_every_row",
)
async def test_find_many_in_and_count_in_read_every_row_past_the_bind_parameter_cap() -> None:
async with _config_table() as (database, _):
values: Final = [*_listed(), *_listed()[:100], "missing"]
rows: Final = await find_many_in(database.litellm_config, "param_name", values)
assert sorted(row.param_name for row in rows) == sorted(_listed())
assert await count_in(database.litellm_config, "param_name", values) == ROWS
assert await count_in(database.litellm_config, "param_name", values, where={"reload_revision": 1}) == 0
@pytest.mark.covers("other.database.chunked_in.update_many_in_updates_every_row_in_a_transaction")
async def test_update_many_in_updates_every_row_inside_one_transaction() -> None:
async with _config_table() as (database, schema):
async with database.tx(timeout=timedelta(seconds=60)) as transaction:
updated: Final = await update_many_in(
transaction.litellm_config,
"param_name",
_listed(),
data={"reload_revision": 7},
atomicity="caller_transaction",
)
assert updated == ROWS
assert _count(schema, sql.SQL("reload_revision = 7 AND param_name LIKE 'listed-%'")) == ROWS
assert _count(schema, sql.SQL("reload_revision = 0 AND param_name LIKE 'outside-%'")) == OUTSIDE
@pytest.mark.covers("other.database.chunked_in.delete_many_in_deletes_every_row")
async def test_delete_many_in_deletes_every_listed_row_and_nothing_else() -> None:
async with _config_table() as (database, schema):
deleted: Final = await delete_many_in(
database.litellm_config, "param_name", _listed(), atomicity="per_chunk_ok", where={"reload_revision": 0}
)
assert deleted == ROWS
assert _count(schema, sql.SQL("TRUE")) == OUTSIDE
@pytest.mark.covers("other.database.chunked_in.key_metadata_recovery_attaches_details_past_the_bind_parameter_cap")
async def test_key_metadata_recovery_attaches_user_details_for_more_users_than_the_bind_parameter_cap() -> None:
async with _user_table(ROWS) as database:
recovered: Final = {f"key-{n}": {"key_alias": f"alias-{n}", "user_id": f"user-{n}"} for n in range(ROWS)}
attached: Final = await attach_user_details(SimpleNamespace(db=database), recovered) # pyright: ignore[reportArgumentType] # only .db is read
assert all(attached[f"key-{n}"].get("user_email") == f"user-{n}@example.com" for n in range(ROWS))

View file

@ -664,3 +664,41 @@ async def test_attach_user_details_claims_no_team_for_a_multi_team_user_session_
assert "team_id" not in attached["cli-session-bob"]
assert attached["cli-session-bob"]["user_email"] == "bob@example.com"
def _user_lookup_by_filter() -> AsyncMock:
async def find_many(*, where):
return [
SimpleNamespace(user_id=user_id, user_email=f"{user_id}@example.com", teams=[])
for user_id in where["user_id"]["in"]
]
return AsyncMock(side_effect=find_many)
@pytest.mark.asyncio
async def test_attach_user_details_chunks_more_than_5000_user_ids_and_merges_every_chunk():
mock_prisma = MagicMock()
mock_prisma.db.litellm_usertable.find_many = _user_lookup_by_filter()
recovered = {f"key-{n}": {"key_alias": f"alias-{n}", "user_id": f"user-{n}"} for n in range(12_001)}
attached = await attach_user_details(mock_prisma, recovered)
sent = [call.kwargs["where"]["user_id"]["in"] for call in mock_prisma.db.litellm_usertable.find_many.call_args_list]
assert [len(chunk) for chunk in sent] == [5_000, 5_000, 2_001]
assert sorted(user_id for chunk in sent for user_id in chunk) == sorted(f"user-{n}" for n in range(12_001))
assert all(attached[f"key-{n}"]["user_email"] == f"user-{n}@example.com" for n in range(12_001))
@pytest.mark.asyncio
async def test_attach_user_details_leaves_metadata_unchanged_when_a_later_chunk_fails():
mock_prisma = MagicMock()
mock_prisma.db.litellm_usertable.find_many = AsyncMock(
side_effect=[[SimpleNamespace(user_id="user-0", user_email="user-0@example.com", teams=[])], PrismaError()]
)
recovered = {f"key-{n}": {"key_alias": f"alias-{n}", "user_id": f"user-{n}"} for n in range(5_001)}
attached = await attach_user_details(mock_prisma, recovered)
assert mock_prisma.db.litellm_usertable.find_many.call_count == 2
assert attached == recovered

View file

@ -0,0 +1,420 @@
"""Tests for tests/code_coverage_tests/check_unbounded_in_lists.py.
The checker reads Python rather than grepping for `IN (`, so the cases that matter are
the ones a grep gets wrong: a subquery or a literal list inside the parentheses, a
runtime value spliced in after them, a fixed display versus a name in a Prisma filter,
and where a `# bounded-ok` marker may sit for a literal a comment cannot go inside.
"""
import importlib.util
import sys
from pathlib import Path
_CHECKER_PATH = Path(__file__).resolve().parents[1] / "code_coverage_tests" / "check_unbounded_in_lists.py"
_SPEC = importlib.util.spec_from_file_location("check_unbounded_in_lists", _CHECKER_PATH)
assert _SPEC is not None and _SPEC.loader is not None
checker = importlib.util.module_from_spec(_SPEC)
sys.modules[_SPEC.name] = checker
_SPEC.loader.exec_module(checker)
def _check(tmp_path: Path, source: str) -> tuple:
target = tmp_path / "module.py"
target.write_text(source, encoding="utf-8")
return checker.check_file(target)
def _kinds(tmp_path: Path, source: str) -> tuple:
return tuple(finding.kind for finding in _check(tmp_path, source))
def _lines(tmp_path: Path, source: str) -> tuple:
return tuple(finding.line for finding in _check(tmp_path, source))
class TestPrismaFilters:
def test_a_name_is_flagged(self, tmp_path):
assert _kinds(tmp_path, 'where = {"user_id": {"in": user_ids}}\n') == ("prisma",)
def test_a_call_is_flagged(self, tmp_path):
assert _kinds(tmp_path, 'where = {"user_id": {"in": list(user_ids)}}\n') == ("prisma",)
def test_a_comprehension_is_flagged(self, tmp_path):
assert _kinds(tmp_path, 'where = {"id": {"in": [row.id for row in rows]}}\n') == ("prisma",)
def test_an_attribute_is_flagged(self, tmp_path):
assert _kinds(tmp_path, 'where = {"user_id": {"in": data.user_ids}}\n') == ("prisma",)
def test_a_starred_display_is_flagged(self, tmp_path):
assert _kinds(tmp_path, 'where = {"user_id": {"in": [*user_ids]}}\n') == ("prisma",)
def test_not_in_is_flagged(self, tmp_path):
assert _kinds(tmp_path, 'where = {"status": {"not_in": list(statuses)}}\n') == ("prisma",)
def test_a_filter_nested_in_a_clause_list_is_flagged(self, tmp_path):
source = 'where = {"OR": [{"team_id": {"in": team_ids}}, {"user_id": user_id}]}\n'
assert _kinds(tmp_path, source) == ("prisma",)
def test_a_display_of_constants_passes(self, tmp_path):
assert _kinds(tmp_path, 'where = {"status": {"not_in": ["failed", "expired"]}}\n') == ()
def test_a_display_with_a_fixed_number_of_names_passes(self, tmp_path):
assert _kinds(tmp_path, 'where = {"user_id": {"in": [user_id]}}\n') == ()
assert _kinds(tmp_path, 'where = {"user_id": {"in": (owner, editor)}}\n') == ()
def test_a_module_constant_bound_to_a_display_passes(self, tmp_path):
constant = 'ANCHORED: Final = frozenset({"oauth2", "api_key"})\n'
assert _kinds(tmp_path, constant + 'where = {"auth_type": {"in": ANCHORED}}\n') == ()
assert _kinds(tmp_path, constant + 'where = {"auth_type": {"in": list(ANCHORED)}}\n') == ()
assert _kinds(tmp_path, constant + 'where = {"auth_type": {"in": sorted(ANCHORED)}}\n') == ()
def test_a_module_constant_built_from_another_passes(self, tmp_path):
source = 'FIRST = ("a", "b")\nSECOND: Final = tuple(FIRST)\nwhere = {"x": {"in": SECOND}}\n'
assert _kinds(tmp_path, source) == ()
def test_casing_does_not_make_a_constant(self, tmp_path):
assert _kinds(tmp_path, 'where = {"auth_type": {"in": ANCHORED_AUTH_TYPES}}\n') == ("prisma",)
assert _kinds(tmp_path, 'USER_IDS = load_ids()\nwhere = {"user_id": {"in": USER_IDS}}\n') == ("prisma",)
assert _kinds(tmp_path, 'from x import STATES\nwhere = {"s": {"in": list(STATES)}}\n') == ("prisma",)
assert _kinds(tmp_path, 'terminal = ("done", "failed")\nwhere = {"s": {"in": terminal}}\n') == ()
def test_a_constant_spread_into_a_display_is_still_a_constant(self, tmp_path):
base = 'BASE: Final = ("a", "b")\n'
assert _kinds(tmp_path, base + 'MORE: Final = (*BASE, "c")\nwhere = {"s": {"not_in": list(MORE)}}\n') == ()
assert _kinds(tmp_path, base + 'where = {"s": {"in": [*BASE, "c"]}}\n') == ()
assert _kinds(tmp_path, base + 'where = {"s": {"in": [*BASE, *extra]}}\n') == ("prisma",)
assert _kinds(tmp_path, 'MORE: Final = (*load(), "c")\nwhere = {"s": {"in": MORE}}\n') == ("prisma",)
def test_a_module_value_that_could_grow_is_not_a_constant(self, tmp_path):
assert _kinds(tmp_path, 'IDS = ["a"]\nIDS.append(late)\nwhere = {"x": {"in": IDS}}\n') == ("prisma",)
assert _kinds(tmp_path, 'IDS = sorted(("a", "b"))\nwhere = {"x": {"in": IDS}}\n') == ("prisma",)
assert _kinds(tmp_path, 'IDS = ("a",)\nwhere = {"x": {"in": IDS}}\n') == ()
assert _kinds(tmp_path, 'IDS = frozenset(["a", "b"])\nwhere = {"x": {"in": IDS}}\n') == ()
def test_an_alias_is_as_fixed_as_what_it_names(self, tmp_path):
assert _kinds(tmp_path, 'A = load_ids()\nB = A\nwhere = {"x": {"in": B}}\n') == ("prisma",)
assert _kinds(tmp_path, 'A = ("a",)\nB = A\nwhere = {"x": {"in": B}}\n') == ()
def test_a_module_name_bound_twice_is_not_a_constant(self, tmp_path):
source = 'IDS = ("a",)\nIDS = load_ids()\nwhere = {"user_id": {"in": IDS}}\n'
assert _kinds(tmp_path, source) == ("prisma",)
def test_a_local_binding_is_not_a_constant(self, tmp_path):
source = 'def f():\n ids = ("a", "b")\n return {"user_id": {"in": ids}}\n'
assert _kinds(tmp_path, source) == ("prisma",)
def test_a_name_wrapped_in_a_constructor_is_still_flagged(self, tmp_path):
assert _kinds(tmp_path, 'where = {"token": {"in": tuple(frozenset(tokens))}}\n') == ("prisma",)
def test_a_scalar_value_passes(self, tmp_path):
assert _kinds(tmp_path, 'parameter = {"name": "q", "in": "query"}\n') == ()
def test_a_dict_with_a_spread_does_not_break_the_walk(self, tmp_path):
assert _kinds(tmp_path, 'where = {**base, "team_id": {"in": team_ids}}\n') == ("prisma",)
def test_the_reported_line_is_the_key_line(self, tmp_path):
source = 'where = {\n "team_id": {\n "in": sorted(team_ids),\n },\n}\n'
assert _lines(tmp_path, source) == (3,)
def test_the_message_names_the_value(self, tmp_path):
(finding,) = _check(tmp_path, 'where = {"user_id": {"in": list(user_ids)}}\n')
assert "list(user_ids)" in finding.message
def test_an_in_list_is_pointed_at_the_chunking_helper(self, tmp_path):
(finding,) = _check(tmp_path, 'where = {"user_id": {"in": user_ids}}\n')
assert "litellm.repositories.chunked_in" in finding.message
def test_a_not_in_list_is_pointed_at_an_array_parameter_since_it_cannot_be_chunked(self, tmp_path):
(finding,) = _check(tmp_path, 'where = {"user_id": {"not_in": user_ids}}\n')
assert "<> ALL($1::text[])" in finding.message
assert "chunked_in" not in finding.message
class TestTypedDictFieldMaps:
"""A functional TypedDict's field map names fields: its "in" key is a type, not a filter."""
def test_a_functional_typed_dict_field_map_is_not_flagged(self, tmp_path):
source = 'Filter = TypedDict("Filter", {"in": NotRequired[Sequence[str]], "notIn": Sequence[str]})\n'
assert _kinds(tmp_path, source) == ()
def test_the_typing_and_typing_extensions_attribute_forms_are_not_flagged(self, tmp_path):
source = (
'A = typing.TypedDict("A", {"in": Sequence[str]})\n'
'B = typing_extensions.TypedDict("B", {"notIn": Sequence[str]})\n'
)
assert _kinds(tmp_path, source) == ()
def test_a_fields_keyword_field_map_is_not_flagged(self, tmp_path):
source = 'Filter = TypedDict("Filter", fields={"in": Sequence[str]}, total=False)\n'
assert _kinds(tmp_path, source) == ()
def test_a_filter_passed_to_another_call_is_still_flagged(self, tmp_path):
source = 'rows = find_many("Filter", {"in": user_ids})\n'
assert _kinds(tmp_path, source) == ("prisma",)
def test_a_typed_dict_from_another_module_is_still_flagged(self, tmp_path):
source = 'Filter = mylib.TypedDict("Filter", {"in": user_ids})\n'
assert _kinds(tmp_path, source) == ("prisma",)
def test_a_filter_nested_inside_a_field_map_value_is_still_flagged(self, tmp_path):
source = 'Filter = TypedDict("Filter", {"where": {"user_id": {"in": user_ids}}})\n'
assert _kinds(tmp_path, source) == ("prisma",)
def test_a_filter_as_the_first_argument_of_typed_dict_is_still_flagged(self, tmp_path):
source = 'Filter = TypedDict({"in": user_ids}, {})\n'
assert _kinds(tmp_path, source) == ("prisma",)
class TestRawSql:
def test_an_fstring_slice_is_flagged(self, tmp_path):
assert _kinds(tmp_path, 'sql = f"WHERE team_id IN ({placeholders})"\n') == ("raw-sql",)
def test_not_in_is_flagged(self, tmp_path):
assert _kinds(tmp_path, "sql = f'\"{field}\" NOT IN ({placeholders})'\n") == ("raw-sql",)
def test_lowercase_sql_is_flagged(self, tmp_path):
assert _kinds(tmp_path, 'sql = f"where team_id in ({placeholders})"\n') == ("raw-sql",)
def test_a_format_slot_is_flagged(self, tmp_path):
assert _kinds(tmp_path, 'sql = "WHERE team_id IN ({})".format(placeholders)\n') == ("raw-sql",)
assert _kinds(tmp_path, 'SQL = "WHERE team_id IN ({ids})"\n') == ("raw-sql",)
def test_a_percent_slot_is_flagged(self, tmp_path):
assert _kinds(tmp_path, 'sql = "WHERE team_id IN (%s)" % placeholders\n') == ("raw-sql",)
assert _kinds(tmp_path, 'sql = "WHERE team_id IN (%(ids)s)" % {"ids": placeholders}\n') == ("raw-sql",)
def test_a_literal_that_closes_after_the_paren_is_flagged(self, tmp_path):
assert _kinds(tmp_path, 'sql = "WHERE team_id IN (" + placeholders + ")"\n') == ("raw-sql",)
def test_a_subquery_passes(self, tmp_path):
source = 'sql = f"""\n DELETE FROM "{table}"\n WHERE id IN (\n SELECT id FROM "{table}" LIMIT $1\n )\n"""\n'
assert _kinds(tmp_path, source) == ()
def test_an_implicitly_concatenated_subquery_passes(self, tmp_path):
source = "sql = (\n 'DELETE FROM t WHERE request_id IN ('\n 'SELECT request_id FROM t LIMIT $1)'\n)\n"
assert _kinds(tmp_path, source) == ()
def test_a_fixed_number_of_placeholders_passes(self, tmp_path):
assert _kinds(tmp_path, 'sql = f"api_key NOT IN (${p}, ${p + 1})"\n') == ()
def test_a_literal_list_passes(self, tmp_path):
assert _kinds(tmp_path, "sql = \"status NOT IN ('failed', 'expired')\"\n") == ()
def test_an_array_parameter_passes(self, tmp_path):
assert _kinds(tmp_path, 'sql = "WHERE user_id = ANY($1::text[])"\n') == ()
assert _kinds(tmp_path, 'sql = "WHERE model IN (SELECT jsonb_array_elements_text($1::jsonb))"\n') == ()
def test_an_escaped_brace_passes(self, tmp_path):
assert _kinds(tmp_path, 'sql = f"WHERE x IN ({{literal}}) AND y = {y}"\n') == ()
def test_a_word_ending_in_in_passes(self, tmp_path):
assert _kinds(tmp_path, 'sql = f"SELECT MIN ({column}) FROM t"\n') == ()
assert _kinds(tmp_path, 'message = f"LOGIN ({user}) failed"\n') == ()
def test_an_fstring_is_reported_once(self, tmp_path):
assert _kinds(tmp_path, 'sql = f"WHERE a IN ({x})" + f" AND b IN ({y})"\n') == ("raw-sql", "raw-sql")
def test_a_multiline_literal_reports_its_first_line_and_names_the_in_line(self, tmp_path):
source = 'sql = f"""\n SELECT 1\n FROM t\n WHERE team_id IN ({placeholders})\n"""\n'
(finding,) = _check(tmp_path, source)
assert finding.line == 1
assert "line 4" in finding.message
class TestMarkers:
def test_a_marker_on_the_line_suppresses(self, tmp_path):
source = 'where = {"team_id": {"in": page_ids}} # bounded-ok: one page of at most 100 ids\n'
assert _kinds(tmp_path, source) == ()
def test_a_marker_shares_the_line_with_other_suppressions(self, tmp_path):
source = 'where = {"team_id": {"in": page_ids}} # mutable-ok: prisma filter # bounded-ok: one page\n'
assert _kinds(tmp_path, source) == ()
def test_a_marker_alone_on_the_line_above_suppresses(self, tmp_path):
source = '# bounded-ok: the expected views are a fixed set\nsql = f"""\n WHERE viewname IN ({views})\n"""\n'
assert _kinds(tmp_path, source) == ()
def test_a_marker_two_lines_above_does_not_suppress(self, tmp_path):
source = '# bounded-ok: one page\n\nwhere = {"team_id": {"in": page_ids}}\n'
assert _kinds(tmp_path, source) == ("prisma",)
def test_a_marker_trailing_the_line_above_does_not_suppress(self, tmp_path):
source = 'other = 1 # bounded-ok: one page\nwhere = {"team_id": {"in": page_ids}}\n'
assert _kinds(tmp_path, source) == ("prisma",)
def test_a_marker_without_a_reason_is_its_own_finding_and_suppresses_nothing(self, tmp_path):
source = 'where = {"team_id": {"in": page_ids}} # bounded-ok\n'
assert _kinds(tmp_path, source) == ("marker", "prisma")
def test_a_marker_with_a_token_reason_is_rejected(self, tmp_path):
source = 'where = {"team_id": {"in": page_ids}} # bounded-ok: ok\n'
assert _kinds(tmp_path, source) == ("marker", "prisma")
class TestDriver:
def test_the_chunking_helper_is_exempt(self):
helper = checker.REPO_ROOT / "litellm" / "repositories" / "chunked_in.py"
assert "prisma" in tuple(finding.kind for finding in checker.check_file(helper))
assert checker.scan(checker.collect_paths([str(helper)])) == ()
def test_a_copy_of_the_helper_elsewhere_is_not_exempt(self, tmp_path):
helper = checker.REPO_ROOT / "litellm" / "repositories" / "chunked_in.py"
copy = tmp_path / "chunked_in.py"
copy.write_text(helper.read_text(encoding="utf-8"), encoding="utf-8")
assert "prisma" in tuple(finding.kind for finding in checker.scan([copy]))
def test_a_syntax_error_is_reported_not_raised(self, tmp_path):
assert _kinds(tmp_path, "def broken(:\n") == ("unreadable",)
def test_directories_are_walked(self, tmp_path):
nested = tmp_path / "pkg" / "sub"
nested.mkdir(parents=True)
(nested / "a.py").write_text('where = {"user_id": {"in": user_ids}}\n', encoding="utf-8")
(nested / "b.txt").write_text('where = {"user_id": {"in": user_ids}}\n', encoding="utf-8")
findings = checker.scan(checker.collect_paths([str(tmp_path / "pkg")]))
assert tuple(finding.path.name for finding in findings) == ("a.py",)
def _identities(tmp_path: Path, source: str) -> tuple:
return tuple(checker.identify(_check(tmp_path, source)))
class TestIdentity:
def test_a_finding_is_keyed_by_scope_field_and_occurrence_not_line(self, tmp_path):
source = (
"class Repo:\n"
" async def load(self):\n"
' a = {"user_id": {"in": ids}}\n'
' b = {"user_id": {"in": more}}\n'
' return {"team_id": {"not_in": teams}}\n'
)
path = (tmp_path / "module.py").resolve().as_posix()
assert _identities(tmp_path, source) == (
f"{path} Repo.load prisma user_id.in `ids` 0",
f"{path} Repo.load prisma user_id.in `more` 0",
f"{path} Repo.load prisma team_id.not_in `teams` 0",
)
def test_the_same_expression_twice_in_a_scope_is_told_apart_by_occurrence(self, tmp_path):
source = 'def f():\n a = {"user_id": {"in": ids}}\n return {"user_id": {"in": ids}}\n'
assert tuple(key.rsplit(" ", 1)[1] for key in _identities(tmp_path, source)) == ("0", "1")
def test_the_value_is_whitespace_normalized(self, tmp_path):
spread = 'def f():\n return {"user_id": {"in": sorted(\n ids ,\n )}}\n'
compact = 'def f():\n return {"user_id": {"in": sorted(ids)}}\n'
assert _identities(tmp_path, spread) == _identities(tmp_path, compact)
def test_the_field_is_read_from_a_subscript_or_keyword_or_computed_key(self, tmp_path):
source = 'where["user_id"] = {"in": ids}\nwhere = Filter(team_id={"in": ids})\nwhere = {field: {"in": ids}}\n'
subjects = tuple(key.split(" ")[3] for key in _identities(tmp_path, source))
assert subjects == ("user_id.in", "team_id.in", "[field].in")
def test_raw_sql_is_keyed_by_the_column_before_in(self, tmp_path):
source = 'def q():\n return f"WHERE \\"{column}\\" NOT IN ({placeholders})"\n'
path = (tmp_path / "module.py").resolve().as_posix()
assert _identities(tmp_path, source) == (f"{path} q raw-sql {{column}}.IN `IN ({{placeholders}})` 0",)
def test_a_raw_sql_value_is_its_normalized_in_slot_without_the_rest_of_the_query(self, tmp_path):
source = 'def q():\n return f"""WHERE id IN (\n {placeholders}\n ) AND deleted = false"""\n'
path = (tmp_path / "module.py").resolve().as_posix()
assert _identities(tmp_path, source) == (f"{path} q raw-sql id.IN `IN ( {{placeholders}} )` 0",)
def test_moving_code_down_the_file_keeps_the_key(self, tmp_path):
source = 'def f():\n return {"user_id": {"in": ids}}\n'
shifted = "import os\n\n\ndef g():\n return 1\n\n\n" + source
assert _identities(tmp_path, source) == _identities(tmp_path, shifted)
class TestReplacedFilter:
"""Swapping a baselined filter for a different unbounded one on the same field must not pass."""
def test_a_replaced_expression_reads_as_one_new_and_one_stale(self, tmp_path, capsys):
target = tmp_path / "module.py"
baseline = tmp_path / "baseline.txt"
target.write_text('def f():\n return {"user_id": {"in": old_ids}}\n', encoding="utf-8")
assert checker.main([str(target), "--baseline", str(baseline), "--update-baseline"]) == 0
target.write_text('def f():\n return {"user_id": {"in": new_ids}}\n', encoding="utf-8")
capsys.readouterr()
assert checker.main([str(target), "--baseline", str(baseline)]) == 1
assert "0 baselined, 1 new, 1 stale" in capsys.readouterr().out
def test_an_identical_expression_re_added_is_the_same_finding(self, tmp_path):
target = tmp_path / "module.py"
baseline = tmp_path / "baseline.txt"
target.write_text('def f():\n return {"user_id": {"in": ids}}\n', encoding="utf-8")
assert checker.main([str(target), "--baseline", str(baseline), "--update-baseline"]) == 0
target.write_text('import os\n\n\ndef f():\n x = 1\n return {"user_id": {"in": ids}}\n', encoding="utf-8")
assert checker.main([str(target), "--baseline", str(baseline)]) == 0
class TestBaseline:
def _run(self, *args: str) -> int:
return checker.main(list(args))
def _write(self, tmp_path: Path, source: str) -> Path:
target = tmp_path / "pkg" / "module.py"
target.parent.mkdir(exist_ok=True)
target.write_text(source, encoding="utf-8")
return target
def test_a_finding_missing_from_the_baseline_fails_the_run(self, tmp_path, capsys):
target = self._write(tmp_path, 'where = {"user_id": {"in": user_ids}}\n')
baseline = tmp_path / "baseline.txt"
assert self._run(str(target), "--baseline", str(baseline)) == 1
out = capsys.readouterr().out
assert f"{target}:1: prisma" in out
assert "1 new" in out
def test_a_baselined_finding_passes_even_after_the_code_moves(self, tmp_path, capsys):
target = self._write(tmp_path, 'def f():\n return {"user_id": {"in": user_ids}}\n')
baseline = tmp_path / "baseline.txt"
assert self._run(str(target), "--baseline", str(baseline), "--update-baseline") == 0
target.write_text("import os\n\n\n" + target.read_text(encoding="utf-8"), encoding="utf-8")
assert self._run(str(target), "--baseline", str(baseline)) == 0
assert "1 baselined, 0 new, 0 stale" in capsys.readouterr().out
def test_a_new_finding_beside_a_baselined_one_fails(self, tmp_path, capsys):
target = self._write(tmp_path, 'def f():\n return {"user_id": {"in": user_ids}}\n')
baseline = tmp_path / "baseline.txt"
assert self._run(str(target), "--baseline", str(baseline), "--update-baseline") == 0
target.write_text(
target.read_text(encoding="utf-8") + 'def g():\n return {"user_id": {"in": user_ids}}\n',
encoding="utf-8",
)
assert self._run(str(target), "--baseline", str(baseline)) == 1
assert f"{target}:4: prisma" in capsys.readouterr().out
def test_a_fixed_finding_leaves_a_stale_entry_that_fails_the_run(self, tmp_path, capsys):
target = self._write(tmp_path, 'def f():\n return {"user_id": {"in": user_ids}}\n')
baseline = tmp_path / "baseline.txt"
assert self._run(str(target), "--baseline", str(baseline), "--update-baseline") == 0
target.write_text('def f():\n return {"user_id": {"in": [user_id]}}\n', encoding="utf-8")
assert self._run(str(target), "--baseline", str(baseline)) == 1
out = capsys.readouterr().out
assert "stale entry" in out
assert "f prisma user_id.in `user_ids` 0" in out
def test_update_baseline_drops_fixed_entries_and_keeps_unscanned_ones(self, tmp_path):
target = self._write(tmp_path, 'def f():\n return {"user_id": {"in": user_ids}}\n')
baseline = tmp_path / "baseline.txt"
elsewhere = "litellm/elsewhere.py g prisma team_id.in 0"
fixed = f"{target.resolve().as_posix()} gone prisma team_id.in 0"
baseline.write_text(f"{elsewhere}\n{fixed}\n", encoding="utf-8")
assert self._run(str(target), "--baseline", str(baseline), "--update-baseline") == 0
assert checker.read_baseline(baseline) == frozenset(
{elsewhere, f"{target.resolve().as_posix()} f prisma user_id.in `user_ids` 0"}
)
assert self._run(str(target), "--baseline", str(baseline)) == 0
def test_entries_for_files_outside_the_scan_are_not_stale(self, tmp_path):
target = self._write(tmp_path, "x = 1\n")
baseline = tmp_path / "baseline.txt"
baseline.write_text("litellm/elsewhere.py g prisma team_id.in 0\n", encoding="utf-8")
assert self._run(str(target), "--baseline", str(baseline)) == 0
def test_an_entry_for_a_deleted_file_under_a_scanned_directory_is_stale(self, tmp_path):
self._write(tmp_path, "x = 1\n")
baseline = tmp_path / "baseline.txt"
gone = (tmp_path / "pkg" / "deleted.py").resolve().as_posix()
baseline.write_text(f"{gone} f prisma user_id.in 0\n", encoding="utf-8")
assert self._run(str(tmp_path / "pkg"), "--baseline", str(baseline)) == 1

View file

@ -0,0 +1,269 @@
from collections.abc import Mapping, Sequence
from dataclasses import dataclass, field
from typing import Final
import pytest
from prisma import models as prisma_models
from prisma.builder import QueryBuilder
from litellm.repositories.chunked_in import (
IN_LIST_CHUNK_SIZE,
MAX_IN_LIST_CHUNK_SIZE,
ChunkedFieldWriteError,
SameFieldFilterError,
count_in,
delete_many_in,
find_many_in,
update_many_in,
)
SIZES: Final = (0, 1, 5_000, 5_001, 12_345)
def _matches(row: Mapping[str, object], where: Mapping[str, object]) -> bool:
def clause(key: str, condition: object) -> bool:
if key == "AND":
return all(_matches(row, part) for part in condition)
if isinstance(condition, Mapping):
return row[key] in condition["in"]
return row[key] == condition
return all(clause(key, condition) for key, condition in where.items())
@dataclass
class FakeTable:
"""Evaluates the filters it is sent against in-memory rows, and records each one."""
rows: list[dict[str, object]]
filters: list[Mapping[str, object]] = field(default_factory=list)
def _select(self, where: Mapping[str, object]) -> list[dict[str, object]]:
self.filters.append(where)
return [row for row in self.rows if _matches(row, where)]
async def find_many(self, *, where: Mapping[str, object]) -> Sequence[dict[str, object]]:
return self._select(where)
async def count(self, *, where: Mapping[str, object]) -> int:
return len(self._select(where))
async def update_many(self, *, data: Mapping[str, object], where: Mapping[str, object]) -> int:
selected = self._select(where)
for row in selected:
row.update(data)
return len(selected)
async def delete_many(self, *, where: Mapping[str, object]) -> int:
selected = self._select(where)
self.rows = [row for row in self.rows if row not in selected]
return len(selected)
def in_list_sizes(self) -> list[int]:
return [len(_membership(where)["in"]) for where in self.filters]
def _membership(where: Mapping[str, object]) -> Mapping[str, Sequence[object]]:
inner = where["AND"][1] if "AND" in where else where
((_, condition),) = inner.items()
return condition
def _table(size: int) -> FakeTable:
return FakeTable(rows=[{"id": f"id-{n}", "team": "even" if n % 2 == 0 else "odd"} for n in range(size + 10)])
def _ids(size: int) -> list[str]:
return [f"id-{n}" for n in range(size)]
def _expected_chunks(size: int, chunk_size: int = IN_LIST_CHUNK_SIZE) -> list[int]:
return [min(chunk_size, size - start) for start in range(0, size, chunk_size)]
@pytest.mark.parametrize("size", SIZES)
async def test_find_many_in_returns_every_matching_row_in_bounded_chunks(size: int) -> None:
table = _table(size)
rows = await find_many_in(table, "id", _ids(size))
assert [row["id"] for row in rows] == _ids(size)
assert table.in_list_sizes() == _expected_chunks(size)
@pytest.mark.parametrize("size", SIZES)
async def test_count_in_sums_the_chunk_counts(size: int) -> None:
table = _table(size)
assert await count_in(table, "id", _ids(size)) == size
assert table.in_list_sizes() == _expected_chunks(size)
@pytest.mark.parametrize("size", SIZES)
async def test_update_many_in_updates_every_row_and_sums_counts(size: int) -> None:
table = _table(size)
updated = await update_many_in(table, "id", _ids(size), data={"team": "moved"}, atomicity="per_chunk_ok")
assert updated == size
assert [row["id"] for row in table.rows if row["team"] == "moved"] == _ids(size)
assert table.in_list_sizes() == _expected_chunks(size)
@pytest.mark.parametrize("size", SIZES)
async def test_delete_many_in_deletes_every_row_and_sums_counts(size: int) -> None:
table = _table(size)
deleted = await delete_many_in(table, "id", _ids(size), atomicity="caller_transaction")
assert deleted == size
assert [row["id"] for row in table.rows] == [f"id-{n}" for n in range(size, size + 10)]
assert table.in_list_sizes() == _expected_chunks(size)
async def test_an_empty_list_sends_no_query() -> None:
table = _table(0)
assert await find_many_in(table, "id", []) == ()
assert await count_in(table, "id", []) == 0
assert await update_many_in(table, "id", [], data={"team": "x"}, atomicity="per_chunk_ok") == 0
assert await delete_many_in(table, "id", [], atomicity="per_chunk_ok") == 0
assert table.filters == []
async def test_duplicate_values_are_sent_once_in_first_seen_order() -> None:
table = _table(IN_LIST_CHUNK_SIZE + 1)
values = [*reversed(_ids(IN_LIST_CHUNK_SIZE + 1)), *_ids(IN_LIST_CHUNK_SIZE + 1)]
assert await count_in(table, "id", values) == IN_LIST_CHUNK_SIZE + 1
sent = [value for where in table.filters for value in _membership(where)["in"]]
assert sent == list(reversed(_ids(IN_LIST_CHUNK_SIZE + 1)))
async def test_where_is_anded_with_each_chunk() -> None:
table = _table(12_345)
where = {"team": "even"}
rows = await find_many_in(table, "id", _ids(12_345), where=where)
assert [row["id"] for row in rows] == [f"id-{n}" for n in range(0, 12_345, 2)]
assert [set(where_sent) for where_sent in table.filters] == [{"AND"}] * 3
assert all(where_sent["AND"][0] == where for where_sent in table.filters)
assert table.in_list_sizes() == _expected_chunks(12_345)
@pytest.mark.parametrize(
"where",
[
{"id": "id-1"},
{"id": {"not": "id-1"}},
{"AND": [{"team": "even"}, {"id": {"in": ["id-1"]}}]},
{"OR": ({"id": "id-1"},)},
{"NOT": {"id": "id-1"}},
{"AND": [{"OR": [{"NOT": {"id": "id-1"}}]}]},
],
)
async def test_where_filtering_the_chunked_field_is_refused_before_any_query(where: Mapping[str, object]) -> None:
table = _table(3)
with pytest.raises(SameFieldFilterError, match="`id`"):
await count_in(table, "id", _ids(3), where=where)
assert table.filters == []
async def test_writes_require_an_atomicity_decision() -> None:
table = _table(1)
with pytest.raises(TypeError, match="atomicity"):
await update_many_in(table, "id", _ids(1), data={"team": "x"}) # pyright: ignore[reportCallIssue] # the missing argument is the test
with pytest.raises(TypeError, match="atomicity"):
await delete_many_in(table, "id", _ids(1)) # pyright: ignore[reportCallIssue] # the missing argument is the test
assert table.filters == []
def _find_many_query(where: Mapping[str, object]) -> str:
return QueryBuilder(
method="find_many", model=prisma_models.LiteLLM_Config, arguments={"where": where}
).build_query()
async def test_the_composed_filter_renders_like_a_hand_written_prisma_filter() -> None:
table = FakeTable(rows=[{"param_name": "a", "param_value": 1}])
await find_many_in(table, "param_name", ["a", "b", "a"], where={"param_value": 1})
hand_written = {"AND": [{"param_value": 1}, {"param_name": {"in": ["a", "b"]}}]}
assert _find_many_query(table.filters[0]) == _find_many_query(hand_written)
async def _run_every_operation(table: FakeTable, values: Sequence[str], chunk_size: int) -> None:
await find_many_in(table, "id", values, chunk_size=chunk_size)
await count_in(table, "id", values, chunk_size=chunk_size)
await update_many_in(table, "id", values, data={"team": "x"}, atomicity="per_chunk_ok", chunk_size=chunk_size)
await delete_many_in(table, "id", values, atomicity="per_chunk_ok", chunk_size=chunk_size)
async def test_the_default_chunk_size_is_unchanged() -> None:
assert IN_LIST_CHUNK_SIZE == 5_000
assert MAX_IN_LIST_CHUNK_SIZE == 30_000
@pytest.mark.parametrize("chunk_size", [7, 100, 1_234])
async def test_a_custom_chunk_size_sets_the_number_of_queries_for_every_operation(chunk_size: int) -> None:
table = _table(1_234)
await _run_every_operation(table, _ids(1_234), chunk_size)
assert table.in_list_sizes() == _expected_chunks(1_234, chunk_size) * 4
assert table.rows == [{"id": f"id-{n}", "team": "even" if n % 2 == 0 else "odd"} for n in range(1_234, 1_244)]
@dataclass
class ChunkSizeRecorder:
"""Counts every value it is sent without scanning rows, so large chunks stay cheap."""
sizes: list[int] = field(default_factory=list)
async def count(self, *, where: Mapping[str, object]) -> int:
self.sizes.append(len(_membership(where)["in"]))
return self.sizes[-1]
@pytest.mark.parametrize(
("chunk_size", "expected"),
[(1, [1] * 5), (MAX_IN_LIST_CHUNK_SIZE, [MAX_IN_LIST_CHUNK_SIZE, 1])],
)
async def test_the_chunk_size_bounds_are_accepted(chunk_size: int, expected: list[int]) -> None:
table = ChunkSizeRecorder()
size = sum(expected)
assert await count_in(table, "id", _ids(size), chunk_size=chunk_size) == size
assert table.sizes == expected
@pytest.mark.parametrize("chunk_size", [-1, 0, MAX_IN_LIST_CHUNK_SIZE + 1])
@pytest.mark.parametrize("values", [[], ["id-0"]])
async def test_a_chunk_size_outside_1_to_the_max_is_refused_before_any_query(
chunk_size: int, values: list[str]
) -> None:
table = _table(1)
operations = (
find_many_in(table, "id", values, chunk_size=chunk_size),
count_in(table, "id", values, chunk_size=chunk_size),
update_many_in(table, "id", values, data={"team": "x"}, atomicity="per_chunk_ok", chunk_size=chunk_size),
delete_many_in(table, "id", values, atomicity="per_chunk_ok", chunk_size=chunk_size),
)
for operation in operations:
with pytest.raises(ValueError, match="chunk_size"):
await operation
assert table.filters == []
async def test_the_chunk_filter_equals_a_hand_written_filter() -> None:
table = _table(2)
await find_many_in(table, "id", ["id-0", "id-1", "id-0"])
assert table.filters == [{"id": {"in": ["id-0", "id-1"]}}]
async def test_an_update_that_moves_a_row_into_a_later_chunk_is_refused_before_any_query() -> None:
table = FakeTable(rows=[{"id": "old", "team": "a"}, {"id": "new", "team": "b"}])
with pytest.raises(ChunkedFieldWriteError, match="`id`"):
await update_many_in(table, "id", ["old", "new"], data={"id": "new"}, atomicity="per_chunk_ok", chunk_size=1)
assert table.filters == []
assert table.rows == [{"id": "old", "team": "a"}, {"id": "new", "team": "b"}]
@pytest.mark.parametrize("data", [{"id": "x"}, {"id": {"set": "x"}}, {"team": "x", "id": None}])
@pytest.mark.parametrize("values", [[], ["id-0"]])
async def test_writing_the_chunked_field_is_refused_in_any_form(data: Mapping[str, object], values: list[str]) -> None:
table = _table(1)
with pytest.raises(ChunkedFieldWriteError):
await update_many_in(table, "id", values, data=data, atomicity="per_chunk_ok")
assert table.filters == []
async def test_writing_another_field_that_names_the_chunked_one_is_allowed() -> None:
table = _table(1)
assert await update_many_in(table, "id", ["id-0"], data={"team": {"set": "id"}}, atomicity="per_chunk_ok") == 1