diff --git a/.github/workflows/test-code-quality.yml b/.github/workflows/test-code-quality.yml index 75f645086fb..23955e33dec 100644 --- a/.github/workflows/test-code-quality.yml +++ b/.github/workflows/test-code-quality.yml @@ -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 diff --git a/litellm/proxy/spend_tracking/key_metadata_recovery.py b/litellm/proxy/spend_tracking/key_metadata_recovery.py index 965cded59c4..ce96dc62780 100644 --- a/litellm/proxy/spend_tracking/key_metadata_recovery.py +++ b/litellm/proxy/spend_tracking/key_metadata_recovery.py @@ -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), ) diff --git a/litellm/repositories/chunked_in.py b/litellm/repositories/chunked_in.py new file mode 100644 index 00000000000..d16cb7c991c --- /dev/null +++ b/litellm/repositories/chunked_in.py @@ -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)) diff --git a/litellm/repositories/prisma_protocols.py b/litellm/repositories/prisma_protocols.py index 60c16fbd746..c42301a9316 100644 --- a/litellm/repositories/prisma_protocols.py +++ b/litellm/repositories/prisma_protocols.py @@ -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: ... diff --git a/tests/code_coverage_tests/check_unbounded_in_lists.py b/tests/code_coverage_tests/check_unbounded_in_lists.py new file mode 100644 index 00000000000..6a1aceed04f --- /dev/null +++ b/tests/code_coverage_tests/check_unbounded_in_lists.py @@ -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: ` 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 = "" +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[^#]*))?") +# 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 ``.""" + 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: `", + 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: `", + 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: `", + 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:])) diff --git a/tests/code_coverage_tests/unbounded_in_baseline.txt b/tests/code_coverage_tests/unbounded_in_baseline.txt new file mode 100644 index 00000000000..b1552d90a91 --- /dev/null +++ b/tests/code_coverage_tests/unbounded_in_baseline.txt @@ -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 raw-sql classifier_type.IN `IN ({_LLM_CLASSIFIER_TYPES_SQL})` 0 diff --git a/tests/integration/database/test_chunked_in_lists.py b/tests/integration/database/test_chunked_in_lists.py new file mode 100644 index 00000000000..7cb3e038479 --- /dev/null +++ b/tests/integration/database/test_chunked_in_lists.py @@ -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)) diff --git a/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py b/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py index 1967d7b6aad..acd03964bf3 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py +++ b/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py @@ -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 diff --git a/tests/test_litellm/test_check_unbounded_in_lists.py b/tests/test_litellm/test_check_unbounded_in_lists.py new file mode 100644 index 00000000000..d4f1c97aca7 --- /dev/null +++ b/tests/test_litellm/test_check_unbounded_in_lists.py @@ -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 diff --git a/tests/unit/repositories/test_chunked_in.py b/tests/unit/repositories/test_chunked_in.py new file mode 100644 index 00000000000..0a6eb3aa39b --- /dev/null +++ b/tests/unit/repositories/test_chunked_in.py @@ -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