mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
140 lines
5.5 KiB
Python
140 lines
5.5 KiB
Python
"""Contract machinery shared by every LiteLLM-defined list route, on any surface."""
|
|
|
|
from collections.abc import Sequence
|
|
from typing import Final
|
|
from urllib.parse import urlencode
|
|
|
|
from fastapi import Request
|
|
from fastapi.dependencies.utils import get_flat_params
|
|
from fastapi.params import ParamTypes
|
|
from fastapi.responses import JSONResponse
|
|
from typing_extensions import ReadOnly, TypedDict
|
|
|
|
from litellm.types.proxy.management_endpoints.management_v1 import (
|
|
ListLinks,
|
|
PageLinks,
|
|
ProblemDetail,
|
|
)
|
|
|
|
PROBLEM_CONTENT_TYPE: Final = "application/problem+json"
|
|
# A URN, not an https URL: RFC 9457 only asks that `type` identify the problem
|
|
# type, and an https URI promises documentation at that address. Switch to an
|
|
# https base only when pages actually exist to serve.
|
|
PROBLEM_TYPE_BASE: Final = "urn:litellm:error:"
|
|
|
|
|
|
class ManagementProblem(Exception):
|
|
"""Raised to return an RFC 9457 problem instead of the proxy's OpenAI error shape."""
|
|
|
|
def __init__(self, problem: ProblemDetail) -> None:
|
|
self.problem = problem
|
|
super().__init__(problem.detail)
|
|
|
|
|
|
def problem_response(problem: ProblemDetail) -> JSONResponse:
|
|
return JSONResponse(
|
|
status_code=problem.status,
|
|
content=problem.model_dump(exclude_none=True),
|
|
media_type=PROBLEM_CONTENT_TYPE,
|
|
)
|
|
|
|
|
|
def _declared_query_params(request: Request) -> frozenset[str]:
|
|
route: Final = request.scope.get("route")
|
|
dependant: Final = getattr(route, "dependant", None)
|
|
if dependant is None:
|
|
return frozenset()
|
|
# fastapi>=0.140.7 removed get_flat_dependant(); get_flat_params() returns the
|
|
# flattened (deduped) param list. Filter to query params to match the old behavior.
|
|
return frozenset(
|
|
field.alias
|
|
for field in get_flat_params(dependant)
|
|
if getattr(field.field_info, "in_", None) == ParamTypes.query
|
|
)
|
|
|
|
|
|
def escape_like(value: str) -> str:
|
|
"""Escape LIKE/ILIKE metacharacters. Ids routinely contain `_`, which is a wildcard unescaped."""
|
|
return value.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_")
|
|
|
|
|
|
class ValidationErrorDetail(TypedDict):
|
|
"""The keys of a pydantic/FastAPI validation error a problem document needs."""
|
|
|
|
type: ReadOnly[str]
|
|
loc: ReadOnly[tuple[int | str, ...]]
|
|
msg: ReadOnly[str]
|
|
|
|
|
|
def _is_length_error_of_rejected_items(error: ValidationErrorDetail, errors: Sequence[ValidationErrorDetail]) -> bool:
|
|
"""pydantic counts only items that validated, so a bad item also trips the parent's min_length."""
|
|
return error["type"] == "too_short" and any(
|
|
len(other["loc"]) > len(error["loc"]) and other["loc"][: len(error["loc"])] == error["loc"] for other in errors
|
|
)
|
|
|
|
|
|
def request_validation_problem(raw_errors: Sequence[ValidationErrorDetail]) -> ProblemDetail:
|
|
"""A body that fails validation (an unknown field included) is 422; a bad query parameter is 400."""
|
|
errors: Final = tuple(error for error in raw_errors if not _is_length_error_of_rejected_items(error, raw_errors))
|
|
detail: Final = "; ".join(f"{'.'.join(str(part) for part in error['loc'][1:])}: {error['msg']}" for error in errors)
|
|
if any(error["loc"] and error["loc"][0] == "body" for error in errors):
|
|
return ProblemDetail(
|
|
type=f"{PROBLEM_TYPE_BASE}invalid-request-body",
|
|
title="Invalid request body",
|
|
status=422,
|
|
detail=detail or "The request body is invalid.",
|
|
)
|
|
return ProblemDetail(
|
|
type=f"{PROBLEM_TYPE_BASE}invalid-query-parameter",
|
|
title="Invalid query parameter",
|
|
status=400,
|
|
detail=detail or "The request query parameters are invalid.",
|
|
)
|
|
|
|
|
|
def unknown_query_param_problem(unknown: tuple[str, ...], allowed: tuple[str, ...]) -> ProblemDetail:
|
|
return ProblemDetail(
|
|
type=f"{PROBLEM_TYPE_BASE}unknown-query-parameter",
|
|
title="Unknown query parameter",
|
|
status=400,
|
|
detail=f"Unrecognized query parameter(s): {', '.join(unknown)}.",
|
|
allowed=sorted(allowed),
|
|
)
|
|
|
|
|
|
async def reject_unknown_query_params(request: Request) -> None:
|
|
"""Reject any query param the route did not declare.
|
|
|
|
A silently ignored filter over-returns data, which is worse than a rejected
|
|
request; a fresh surface is the only chance to be strict about it.
|
|
"""
|
|
declared: Final = _declared_query_params(request)
|
|
unknown: Final[tuple[str, ...]] = tuple(sorted(name for name in request.query_params if name not in declared))
|
|
if not unknown:
|
|
return
|
|
raise ManagementProblem(unknown_query_param_problem(unknown=unknown, allowed=tuple(sorted(declared))))
|
|
|
|
|
|
def _page_url(request: Request, page: int) -> str:
|
|
others: Final = tuple((key, value) for key, value in request.query_params.multi_items() if key != "page")
|
|
return f"{request.url.path}?{urlencode((*others, ('page', page)))}"
|
|
|
|
|
|
def build_page_links(request: Request, page: int, has_more: bool) -> PageLinks:
|
|
return PageLinks(
|
|
self_link=_page_url(request, page),
|
|
prev=_page_url(request, page - 1) if page > 1 else None,
|
|
next=_page_url(request, page + 1) if has_more else None,
|
|
)
|
|
|
|
|
|
def build_list_links(request: Request, page: int, total_pages: int) -> ListLinks:
|
|
"""Page-mode links. `last` clamps to page 1 on an empty result set so every link still resolves."""
|
|
last: Final = max(total_pages, 1)
|
|
return ListLinks(
|
|
self_link=_page_url(request, page),
|
|
first=_page_url(request, 1),
|
|
prev=_page_url(request, page - 1) if page > 1 else None,
|
|
next=_page_url(request, page + 1) if page < last else None,
|
|
last=_page_url(request, last),
|
|
)
|