Reduce basedpyright Any counts by ~5.5k across 12 hot files

Types the highest-volume Any sources in the proxy and core paths:
Prisma row seams, **row.model_dump() unpacks, response payloads, and
untyped helper returns.

reportAny        26979 -> 21998  (-4981)
reportExplicitAny  7432 ->  6918   (-514)
total errors    157609 -> 151413  (-6196)

Ratchets basedpyright-code-budget.json down by 6215 across 48 rules so
the gains cannot be given back.
This commit is contained in:
mateo-berri 2026-07-24 19:53:01 -07:00
parent 76b0b10908
commit 6afef76268
13 changed files with 3259 additions and 1217 deletions

View file

@ -1,6 +1,6 @@
{
"reportAny": {
"limit": 37484
"limit": 32503
},
"reportArgumentType": {
"limit": 2704
@ -24,7 +24,7 @@
"limit": 42
},
"reportExplicitAny": {
"limit": 10389
"limit": 9875
},
"reportFunctionMemberAccess": {
"limit": 11
@ -54,10 +54,10 @@
"limit": 0
},
"reportMissingParameterType": {
"limit": 5900
"limit": 5846
},
"reportMissingTypeArgument": {
"limit": 15903
"limit": 15857
},
"reportMissingTypeStubs": {
"limit": 41
@ -99,19 +99,19 @@
"limit": 0
},
"reportUnknownArgumentType": {
"limit": 45894
"limit": 45607
},
"reportUnknownLambdaType": {
"limit": 113
},
"reportUnknownMemberType": {
"limit": 40539
"limit": 40441
},
"reportUnknownParameterType": {
"limit": 20403
"limit": 20307
},
"reportUnknownVariableType": {
"limit": 32141
"limit": 32003
},
"reportUnnecessaryCast": {
"limit": 177
@ -138,7 +138,7 @@
"limit": 206
},
"reportUnusedImport": {
"limit": 1005
"limit": 1004
},
"reportUnusedVariable": {
"limit": 1297

View file

@ -14,8 +14,11 @@ from typing import (
Dict,
Iterator,
List,
Mapping,
NoReturn,
Optional,
Protocol,
Sequence,
Union,
cast,
)
@ -23,6 +26,7 @@ from typing import (
import anyio
import httpx
from pydantic import BaseModel
from typing_extensions import Required, TypedDict
import litellm
from litellm import verbose_logger
@ -60,10 +64,71 @@ FUNCTION_CALL_ATTRIBUTE = "function_call"
_SYNC_ITER_EXHAUSTED = object()
_GCHUNK_FIELDS: frozenset = frozenset(GChunk.__annotations__)
_GCHUNK_FIELDS: frozenset[str] = frozenset(GChunk.__annotations__)
def _next_sync_or_exhausted(it: Any) -> Any:
class _LoggedLiteLLMParams(TypedDict, total=False):
merge_reasoning_content_in_choices: bool | None
class _DumpedDeltaKwargs(TypedDict, total=False):
role: str | None
finish_reason: str | None
tool_calls: Sequence[Mapping[str, object]] | None
class _ModelResponseStreamInitKwargs(TypedDict, total=False):
model: str | None
class _PredibaseStreamDetails(TypedDict):
finish_reason: str
class _PredibaseStreamData(TypedDict, total=False):
token: Mapping[str, str]
details: Required[_PredibaseStreamDetails]
generated_text: str
error: str
class _Ai21StreamData(TypedDict):
completions: Sequence[Mapping[str, Mapping[str, str]]]
class _MaritalkStreamData(TypedDict):
answer: str
class _NlpCloudStreamData(TypedDict):
generated_text: str
class _AlephAlphaStreamData(TypedDict):
completions: Sequence[Mapping[str, str]]
class _AzureStreamChoice(TypedDict):
delta: Mapping[str, str] | None
finish_reason: str | None
class _AzureStreamData(TypedDict):
choices: Sequence[_AzureStreamChoice]
class _BasetenStreamData(TypedDict, total=False):
token: Mapping[str, str]
model_output: Mapping[str, Union[Sequence[str], str]] | str
completion: object
class _TextCompletionStreamChoice(Protocol):
text: str
finish_reason: str | None
def _next_sync_or_exhausted(it: Iterator[object]) -> object:
"""
Call next(it) from a thread and return _SYNC_ITER_EXHAUSTED on StopIteration.
@ -77,7 +142,7 @@ def _next_sync_or_exhausted(it: Any) -> Any:
return _SYNC_ITER_EXHAUSTED
def is_async_iterable(obj: Any) -> bool:
def is_async_iterable(obj: object) -> bool:
"""
Check if an object is an async iterable (can be used with 'async for').
@ -90,7 +155,7 @@ def is_async_iterable(obj: Any) -> bool:
return isinstance(obj, collections.abc.AsyncIterable)
def print_verbose(print_statement):
def print_verbose(print_statement: object):
try:
if litellm.set_verbose:
print(print_statement) # noqa: T201
@ -116,7 +181,7 @@ class CustomStreamWrapper:
self,
completion_stream,
model,
logging_obj: Any,
logging_obj: LiteLLMLoggingObject,
custom_llm_provider: Optional[str] = None,
stream_options=None,
make_call: Optional[Callable] = None,
@ -131,9 +196,8 @@ class CustomStreamWrapper:
self.sent_last_chunk = False
self._stream_created_time: float = time.time()
litellm_params: GenericLiteLLMParams = GenericLiteLLMParams(
**self.logging_obj.model_call_details.get("litellm_params", {})
)
_logged_litellm_params: _LoggedLiteLLMParams = self.logging_obj.model_call_details.get("litellm_params", {})
litellm_params: GenericLiteLLMParams = GenericLiteLLMParams(**_logged_litellm_params)
self.merge_reasoning_content_in_choices: bool = litellm_params.merge_reasoning_content_in_choices or False
self.sent_first_thinking_block = False
self.sent_last_thinking_block = False
@ -195,7 +259,7 @@ class CustomStreamWrapper:
# Snapshot assumes self._hidden_params is populated from litellm_params
# at init and never mutated during the stream. If that ever changes,
# this cache must be removed.
self._base_hidden_params: Dict[str, Any] = {
self._base_hidden_params: dict[str, object] = {
**self._hidden_params,
"response_cost": None,
}
@ -257,7 +321,7 @@ class CustomStreamWrapper:
return False
def process_chunk(self, chunk: str):
def process_chunk(self, chunk: str) -> str:
"""
NLP Cloud streaming returns the entire response, for each chunk. Process this, to only return the delta.
"""
@ -358,7 +422,7 @@ class CustomStreamWrapper:
finish_reason = ""
print_verbose(f"chunk: {chunk}")
if chunk.startswith("data:"):
data_json = json.loads(chunk[5:])
data_json: _PredibaseStreamData = json.loads(chunk[5:])
print_verbose(f"data json: {data_json}")
if "token" in data_json and "text" in data_json["token"]:
text = data_json["token"]["text"]
@ -388,7 +452,7 @@ class CustomStreamWrapper:
def handle_ai21_chunk(self, chunk): # fake streaming
chunk = chunk.decode("utf-8")
data_json = json.loads(chunk)
data_json: _Ai21StreamData = json.loads(chunk)
try:
text = data_json["completions"][0]["data"]["text"]
is_finished = True
@ -403,7 +467,7 @@ class CustomStreamWrapper:
def handle_maritalk_chunk(self, chunk): # fake streaming
chunk = chunk.decode("utf-8")
data_json = json.loads(chunk)
data_json: _MaritalkStreamData = json.loads(chunk)
try:
text = data_json["answer"]
is_finished = True
@ -424,7 +488,7 @@ class CustomStreamWrapper:
if self.model and "dolphin" in self.model:
chunk = self.process_chunk(chunk=chunk)
else:
data_json = json.loads(chunk)
data_json: _NlpCloudStreamData = json.loads(chunk)
chunk = data_json["generated_text"]
text = chunk
if "[DONE]" in text:
@ -441,7 +505,7 @@ class CustomStreamWrapper:
def handle_aleph_alpha_chunk(self, chunk):
chunk = chunk.decode("utf-8")
data_json = json.loads(chunk)
data_json: _AlephAlphaStreamData = json.loads(chunk)
try:
text = data_json["completions"][0]["completion"]
is_finished = True
@ -454,7 +518,7 @@ class CustomStreamWrapper:
except Exception:
raise ValueError(f"Unable to parse response. Original response: {chunk}")
def handle_azure_chunk(self, chunk):
def handle_azure_chunk(self, chunk: str):
is_finished = False
finish_reason = ""
text = ""
@ -469,7 +533,7 @@ class CustomStreamWrapper:
"finish_reason": finish_reason,
}
elif chunk.startswith("data:"):
data_json = json.loads(chunk[5:]) # chunk.startswith("data:"):
data_json: _AzureStreamData = json.loads(chunk[5:]) # chunk.startswith("data:"):
try:
if len(data_json["choices"]) > 0:
delta = data_json["choices"][0]["delta"]
@ -558,7 +622,7 @@ class CustomStreamWrapper:
text = ""
is_finished = False
finish_reason = None
choices = getattr(chunk, "choices", [])
choices: Sequence[_TextCompletionStreamChoice] = getattr(chunk, "choices", [])
if len(choices) > 0:
text = choices[0].text
if choices[0].finish_reason is not None:
@ -579,7 +643,7 @@ class CustomStreamWrapper:
is_finished = False
finish_reason = None
usage = None
choices = getattr(chunk, "choices", [])
choices: Sequence[_TextCompletionStreamChoice] = getattr(chunk, "choices", [])
if len(choices) > 0:
text = choices[0].text
if choices[0].finish_reason is not None:
@ -601,7 +665,7 @@ class CustomStreamWrapper:
chunk = chunk.decode("utf-8")
if len(chunk) > 0:
if chunk.startswith("data:"):
data_json = json.loads(chunk[5:])
data_json: _BasetenStreamData = json.loads(chunk[5:])
if "token" in data_json and "text" in data_json["token"]:
return data_json["token"]["text"]
else:
@ -665,19 +729,23 @@ class CustomStreamWrapper:
except Exception as e:
raise e
def model_response_creator(self, chunk: Optional[dict] = None, hidden_params: Optional[dict] = None):
def model_response_creator(self, chunk: dict[str, object] | None = None, hidden_params: dict | None = None):
_model = self._cached_model_name
_logging_obj_llm_provider = self._cached_logging_llm_provider
if chunk is None:
args: Dict[str, Any] = {"model": _model}
args: dict[str, object] = {"model": _model}
else:
chunk.pop("model", None)
args = {"model": _model}
if chunk:
args.update({k: v for k, v in chunk.items() if k != "stream"})
model_response = ModelResponseStream(**args)
model_response = ModelResponseStream(
**cast( # cast-ok: heterogeneous chunk payload consumed via ModelResponseStream(**kwargs)
_ModelResponseStreamInitKwargs, args
)
)
if self.response_id is not None:
model_response.id = self.response_id
if self.system_fingerprint is not None:
@ -750,7 +818,11 @@ class CustomStreamWrapper:
"""
Copy provider_specific_fields from original_chunk to model_response.
"""
provider_specific_fields = getattr(original_chunk, "provider_specific_fields", None)
provider_specific_fields: dict[str, object] | None = (
getattr( # mutable-ok: assigned straight to ModelResponseStream.provider_specific_fields, declared Optional[Dict[str, Any]]
original_chunk, "provider_specific_fields", None
)
)
if provider_specific_fields is not None:
model_response.provider_specific_fields = provider_specific_fields
for k, v in provider_specific_fields.items():
@ -814,7 +886,9 @@ class CustomStreamWrapper:
model_response.choices[0].delta["role"] = "assistant"
self.sent_first_chunk = True
elif self.sent_first_chunk is True and hasattr(model_response.choices[0].delta, "role"):
_initial_delta = model_response.choices[0].delta.model_dump()
_initial_delta = cast( # cast-ok: Delta.model_dump payload re-enters Delta(**kwargs) which takes extras
_DumpedDeltaKwargs, model_response.choices[0].delta.model_dump()
)
_initial_delta.pop("role", None)
model_response.choices[0].delta = Delta(**_initial_delta)
@ -904,14 +978,17 @@ class CustomStreamWrapper:
if hold is False:
## check if openai/azure chunk
original_chunk = response_obj.get("original_chunk", None)
original_chunk: ModelResponseStream | None = response_obj.get("original_chunk", None)
if original_chunk:
if len(original_chunk.choices) > 0:
choices = []
for choice in original_chunk.choices:
try:
if isinstance(choice, BaseModel):
choice_json = choice.model_dump() # type: ignore
choice_json = cast( # cast-ok: choice dump feeds StreamingChoices(**kwargs) which takes extras
_DumpedDeltaKwargs,
choice.model_dump(),
)
choice_json.pop(
"finish_reason", None
) # for mistral etc. which return a value in their last chunk (not-openai compatible).
@ -946,7 +1023,11 @@ class CustomStreamWrapper:
self.sent_first_chunk = True
if response_obj.get("provider_specific_fields") is not None:
completion_obj["provider_specific_fields"] = response_obj["provider_specific_fields"]
model_response.choices[0].delta = Delta(**completion_obj)
model_response.choices[0].delta = Delta(
**cast( # cast-ok: completion payload feeds Delta(**kwargs) which takes extras
_DumpedDeltaKwargs, completion_obj
)
)
_index: Optional[int] = completion_obj.get("index")
if _index is not None:
model_response.choices[0].index = _index
@ -1111,7 +1192,9 @@ class CustomStreamWrapper:
for key, value in anthropic_response_obj["provider_specific_fields"].items():
setattr(model_response, key, value)
response_obj = cast(dict[str, Any], anthropic_response_obj)
response_obj = cast( # cast-ok: GenericStreamingChunk narrows to the plain response mapping
dict[str, object], anthropic_response_obj
)
elif self.model == "replicate" or self.custom_llm_provider == "replicate":
response_obj = self.handle_replicate_chunk(chunk)
completion_obj["content"] = response_obj["text"]
@ -1295,15 +1378,19 @@ class CustomStreamWrapper:
if response_obj["is_finished"]:
self.received_finish_reason = response_obj["finish_reason"]
elif self.custom_llm_provider == "cached_response":
chunk = cast(ModelResponseStream, chunk)
chunk_finish_reason = chunk.choices[0].finish_reason
cached_chunk = cast( # cast-ok: cached_response streams replay ModelResponseStream chunks
ModelResponseStream, chunk
)
chunk_finish_reason = cached_chunk.choices[0].finish_reason
response_obj = {
"text": chunk.choices[0].delta.content,
"text": cached_chunk.choices[0].delta.content,
"is_finished": chunk_finish_reason is not None,
"finish_reason": chunk_finish_reason,
"original_chunk": chunk,
"original_chunk": cached_chunk,
"tool_calls": (
chunk.choices[0].delta.tool_calls if hasattr(chunk.choices[0].delta, "tool_calls") else None
cached_chunk.choices[0].delta.tool_calls
if hasattr(cached_chunk.choices[0].delta, "tool_calls")
else None
),
}
@ -1311,11 +1398,11 @@ class CustomStreamWrapper:
if response_obj["tool_calls"] is not None:
completion_obj["tool_calls"] = response_obj["tool_calls"]
print_verbose(f"completion obj content: {completion_obj['content']}")
if hasattr(chunk, "id"):
model_response.id = chunk.id
self.response_id = chunk.id
if hasattr(chunk, "system_fingerprint"):
self.system_fingerprint = chunk.system_fingerprint
if hasattr(cached_chunk, "id"):
model_response.id = cached_chunk.id
self.response_id = cached_chunk.id
if hasattr(cached_chunk, "system_fingerprint"):
self.system_fingerprint = cached_chunk.system_fingerprint
if response_obj["is_finished"]:
self.received_finish_reason = response_obj["finish_reason"]
else: # openai / azure chat model
@ -1392,7 +1479,9 @@ class CustomStreamWrapper:
model_response.model = self.model
## FUNCTION CALL PARSING
original_chunk = response_obj.get("original_chunk") if response_obj is not None else None
original_chunk: ModelResponseStream | None = (
response_obj.get("original_chunk") if response_obj is not None else None
)
if (
original_chunk is not None
): # function / tool calling branch - only set for openai/azure compatible endpoints
@ -1430,7 +1519,9 @@ class CustomStreamWrapper:
is None
):
t.function.arguments = ""
_json_delta = delta.model_dump()
_json_delta = cast( # cast-ok: delta dump re-enters Delta(**kwargs) which takes extras
_DumpedDeltaKwargs, delta.model_dump()
)
if "role" not in _json_delta or _json_delta["role"] is None:
_json_delta["role"] = "assistant" # mistral's api returns role as None
if "tool_calls" in _json_delta and isinstance(_json_delta["tool_calls"], list):
@ -1458,7 +1549,11 @@ class CustomStreamWrapper:
if original_chunk.choices[0].delta is None
else dict(original_chunk.choices[0].delta)
)
model_response.choices[0].delta = Delta(**delta)
model_response.choices[0].delta = Delta(
**cast( # cast-ok: raw delta mapping feeds Delta(**kwargs) which takes extras
_DumpedDeltaKwargs, delta
)
)
except Exception:
model_response.choices[0].delta = Delta()
else:
@ -1497,7 +1592,7 @@ class CustomStreamWrapper:
original_exception=e,
)
def set_logging_event_loop(self, loop):
def set_logging_event_loop(self, loop: asyncio.AbstractEventLoop):
"""
import litellm, asyncio
@ -1672,7 +1767,9 @@ class CustomStreamWrapper:
else:
asyncio.run(self.logging_obj.async_success_handler(processed_chunk, None, None, cache_hit))
## SYNC LOGGING — only for sync SDK entrypoints; async proxy paths export via async_success_handler
litellm_params = self.logging_obj.model_call_details.get("litellm_params", {})
litellm_params: dict[str, object] = self.logging_obj.model_call_details.get(
"litellm_params", {}
) # mutable-ok: passed to Logging._is_sync_litellm_request, whose parameter is typed dict
if self.logging_obj._is_sync_litellm_request(litellm_params):
self.logging_obj.success_handler(processed_chunk, None, None, cache_hit)
@ -1700,7 +1797,9 @@ class CustomStreamWrapper:
usage.cost, copy it into _hidden_params so litellm's cost
calculator uses it instead of a token-based estimate.
"""
_usage = getattr(response, "usage", None)
_usage = cast( # cast-ok: assembled ModelResponse carries an Optional[Usage] attribute
Usage | None, getattr(response, "usage", None)
)
if _usage is not None and hasattr(_usage, "cost") and _usage.cost is not None:
if "additional_headers" not in response._hidden_params:
response._hidden_params["additional_headers"] = {}
@ -2254,7 +2353,7 @@ class CustomStreamWrapper:
return chunk
def calculate_total_usage(chunks: List[ModelResponse]) -> Usage:
def calculate_total_usage(chunks: Sequence[ModelResponse]) -> Usage:
"""Assume most recent usage chunk has total usage uptil then."""
prompt_tokens: int = 0
completion_tokens: int = 0
@ -2287,7 +2386,7 @@ def calculate_total_usage(chunks: List[ModelResponse]) -> Usage:
return returned_usage_chunk
def generic_chunk_has_all_required_fields(chunk: dict) -> bool:
def generic_chunk_has_all_required_fields(chunk: Mapping[str, object]) -> bool:
"""
Checks if the provided chunk dictionary contains all required fields for GenericStreamingChunk.

File diff suppressed because it is too large Load diff

View file

@ -61,6 +61,116 @@ from .common_utils import (
drop_params_from_unprocessable_entity_error,
)
if TYPE_CHECKING:
from openai.types.batch import Errors as BatchErrors
from openai.types.batch_request_counts import BatchRequestCounts
from openai.types.beta.threads.message import Attachment as MessageAttachment
from openai.types.beta.threads.message import (
IncompleteDetails as MessageIncompleteDetails,
)
from typing_extensions import NotRequired, TypedDict
from litellm.types.utils import Usage as LiteLLMUsage
class _FileObjectDump(TypedDict):
id: str
bytes: int
created_at: int
filename: str
object: Literal["file"]
purpose: OpenAIFilesPurpose
status: Literal["uploaded", "processed", "error"]
expires_at: int | None
status_details: str | None
class _BatchDump(TypedDict):
id: str
completion_window: str
created_at: int
endpoint: str
input_file_id: str
object: Literal["batch"]
status: Literal[
"validating",
"failed",
"in_progress",
"finalizing",
"completed",
"expired",
"cancelling",
"cancelled",
]
cancelled_at: int | None
cancelling_at: int | None
completed_at: int | None
error_file_id: str | None
errors: BatchErrors | None
expired_at: int | None
expires_at: int | None
failed_at: int | None
finalizing_at: int | None
in_progress_at: int | None
metadata: Optional[
Dict[str, str]
] # mutable-ok: unpacked into LiteLLMBatch, whose metadata field is declared Dict[str, str]
model: str | None
output_file_id: str | None
request_counts: BatchRequestCounts | None
usage: LiteLLMUsage | None
class _MessageDump(TypedDict):
id: str
assistant_id: str | None
attachments: Optional[
list[MessageAttachment]
] # mutable-ok: unpacked into OpenAIMessage, whose attachments field is declared List[Attachment]
completed_at: int | None
content: list[
MessageContent
] # mutable-ok: unpacked into OpenAIMessage, whose content field is declared List[MessageContent]
created_at: int
incomplete_at: int | None
incomplete_details: MessageIncompleteDetails | None
metadata: Optional[
Dict[str, str]
] # mutable-ok: unpacked into OpenAIMessage, whose metadata field is declared Dict[str, str]
object: Literal["thread.message"]
role: Literal["user", "assistant"]
run_id: str | None
status: Literal["in_progress", "incomplete", "completed"]
thread_id: str
class _ThreadDump(TypedDict):
id: str
created_at: int
metadata: object | None
object: Literal["thread"]
class _RunThreadStreamData(TypedDict):
thread_id: str
assistant_id: str
additional_instructions: str | None
instructions: str | None
metadata: Optional[
Dict[str, str]
] # mutable-ok: unpacked into runs.stream(), whose metadata param is declared Dict[str, str]
model: str | None
tools: Iterable[AssistantToolParam] | None
event_handler: NotRequired[AssistantEventHandler]
class _AsyncRunThreadStreamData(TypedDict):
thread_id: str
assistant_id: str
additional_instructions: str | None
instructions: str | None
metadata: Optional[
Dict[str, str]
] # mutable-ok: unpacked into runs.stream(), whose metadata param is declared Dict[str, str]
model: str | None
tools: Iterable[AssistantToolParam] | None
event_handler: NotRequired[AsyncAssistantEventHandler]
openaiOSeriesConfig = OpenAIOSeriesConfig()
openAIGPT5Config = OpenAIGPT5Config()
@ -73,7 +183,9 @@ class MistralEmbeddingConfig:
def __init__(
self,
) -> None:
locals_ = locals().copy()
locals_: Dict[str, object] = (
locals().copy()
) # mutable-ok: invariant dict pins locals()'s values to object; Mapping re-widens them to Any
for key, value in locals_.items():
if key != "self" and value is not None:
setattr(self.__class__, key, value)
@ -165,7 +277,9 @@ class OpenAIConfig(BaseConfig):
top_p: Optional[int] = None,
response_format: Optional[dict] = None,
) -> None:
locals_ = locals().copy()
locals_: Dict[str, object] = (
locals().copy()
) # mutable-ok: invariant dict pins locals()'s values to object; Mapping re-widens them to Any
for key, value in locals_.items():
if key != "self" and value is not None:
setattr(self.__class__, key, value)
@ -275,16 +389,19 @@ class OpenAIConfig(BaseConfig):
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
encoding: Any,
encoding: object,
api_key: Optional[str] = None,
json_mode: Optional[bool] = None,
) -> ModelResponse:
logging_obj.post_call(original_response=raw_response.text)
logging_obj.model_call_details["response_headers"] = raw_response.headers
raw_response_json = cast( # cast-ok: chat completion responses parse to a JSON object
"Dict[str, object]", raw_response.json()
)
final_response_obj = cast(
ModelResponse,
convert_to_model_response_object(
response_object=raw_response.json(),
response_object=raw_response_json,
model_response_object=model_response,
hidden_params={"headers": raw_response.headers},
_response_headers=dict(raw_response.headers),
@ -313,7 +430,7 @@ class OpenAIConfig(BaseConfig):
streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse],
sync_stream: bool,
json_mode: Optional[bool] = False,
) -> Any:
) -> "OpenAIChatCompletionResponseIterator":
return OpenAIChatCompletionResponseIterator(
streaming_response=streaming_response,
sync_stream=sync_stream,
@ -488,14 +605,14 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
async def _call_agentic_completion_hooks_openai(
self,
response: Any,
response: object,
model: str,
messages: List[Dict],
optional_params: Dict,
logging_obj: LiteLLMLoggingObj,
stream: bool,
litellm_params: Dict,
) -> Optional[Any]:
) -> ModelResponse | None:
"""
Call agentic completion hooks for all custom loggers (OpenAI Chat Completions API).
@ -546,15 +663,18 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
kwargs_with_provider["custom_llm_provider"] = custom_llm_provider
# For OpenAI Chat Completions, use the chat completion agentic loop method
agentic_response = await callback.async_run_chat_completion_agentic_loop(
tools=tool_calls,
model=model,
messages=messages,
response=response,
optional_params=optional_params,
logging_obj=logging_obj,
stream=stream,
kwargs=kwargs_with_provider,
agentic_response = cast( # cast-ok: chat completion agentic loop yields a ModelResponse
"ModelResponse | None",
await callback.async_run_chat_completion_agentic_loop(
tools=tool_calls,
model=model,
messages=messages,
response=response,
optional_params=optional_params,
logging_obj=logging_obj,
stream=stream,
kwargs=kwargs_with_provider,
),
)
# First hook that runs agentic loop wins
return agentic_response
@ -816,7 +936,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
status_code = getattr(e, "status_code", 500)
error_headers = getattr(e, "headers", None)
error_text = getattr(e, "text", str(e))
error_response = getattr(e, "response", None)
error_response: httpx.Response | None = getattr(e, "response", None)
error_body = getattr(e, "body", None)
if error_headers is None and error_response:
error_headers = getattr(error_response, "headers", None)
@ -935,7 +1055,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
raise e
# e.message
except Exception as e:
exception_response = getattr(e, "response", None)
exception_response: httpx.Response | None = getattr(e, "response", None)
status_code = getattr(e, "status_code", 500)
exception_body = getattr(e, "body", None)
error_headers = getattr(e, "headers", None)
@ -1092,7 +1212,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
error_headers = getattr(e, "headers", None)
status_code = getattr(e, "status_code", 500)
error_response = getattr(e, "response", None)
error_response: httpx.Response | None = getattr(e, "response", None)
exception_body = getattr(e, "body", None)
if error_headers is None and error_response:
error_headers = getattr(error_response, "headers", None)
@ -1247,7 +1367,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
status_code = getattr(e, "status_code", 500)
error_headers = getattr(e, "headers", None)
error_text = getattr(e, "text", str(e))
error_response = getattr(e, "response", None)
error_response: httpx.Response | None = getattr(e, "response", None)
if error_headers is None and error_response:
error_headers = getattr(error_response, "headers", None)
raise OpenAIError(status_code=status_code, message=error_text, headers=error_headers)
@ -1333,7 +1453,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
status_code = getattr(e, "status_code", 500)
error_headers = getattr(e, "headers", None)
error_text = getattr(e, "text", str(e))
error_response = getattr(e, "response", None)
error_response: httpx.Response | None = getattr(e, "response", None)
if error_headers is None and error_response:
error_headers = getattr(error_response, "headers", None)
raise OpenAIError(status_code=status_code, message=error_text, headers=error_headers)
@ -1626,7 +1746,10 @@ class OpenAIFilesAPI(BaseLLM):
openai_client: AsyncOpenAI,
) -> OpenAIFileObject:
response = await openai_client.files.create(**create_file_data) # type: ignore[arg-type]
return OpenAIFileObject(**response.model_dump())
response_dict = cast( # cast-ok: FileObject.model_dump returns its typed fields
"_FileObjectDump", response.model_dump()
)
return OpenAIFileObject(**response_dict)
def create_file(
self,
@ -1638,7 +1761,7 @@ class OpenAIFilesAPI(BaseLLM):
max_retries: Optional[int],
organization: Optional[str],
client: Optional[Union[OpenAI, AsyncOpenAI]] = None,
) -> Union[OpenAIFileObject, Coroutine[Any, Any, OpenAIFileObject]]:
) -> OpenAIFileObject | Coroutine[None, None, OpenAIFileObject]:
openai_client: Optional[Union[OpenAI, AsyncOpenAI]] = self.get_openai_client(
api_key=api_key,
api_base=api_base,
@ -1662,7 +1785,10 @@ class OpenAIFilesAPI(BaseLLM):
create_file_data=create_file_data, openai_client=openai_client
)
response = cast(OpenAI, openai_client).files.create(**create_file_data) # type: ignore[arg-type]
return OpenAIFileObject(**response.model_dump())
response_dict = cast( # cast-ok: FileObject.model_dump returns its typed fields
"_FileObjectDump", response.model_dump()
)
return OpenAIFileObject(**response_dict)
async def afile_content(
self,
@ -1682,7 +1808,7 @@ class OpenAIFilesAPI(BaseLLM):
max_retries: Optional[int],
organization: Optional[str],
client: Optional[Union[OpenAI, AsyncOpenAI]] = None,
) -> Union[HttpxBinaryResponseContent, Coroutine[Any, Any, HttpxBinaryResponseContent]]:
) -> HttpxBinaryResponseContent | Coroutine[None, None, HttpxBinaryResponseContent]:
openai_client: Optional[Union[OpenAI, AsyncOpenAI]] = self.get_openai_client(
api_key=api_key,
api_base=api_base,
@ -1986,7 +2112,8 @@ class OpenAIBatchesAPI(BaseLLM):
openai_client: AsyncOpenAI,
) -> LiteLLMBatch:
response = await openai_client.batches.create(**create_batch_data) # type: ignore[arg-type]
return LiteLLMBatch(**response.model_dump())
batch_dict = cast("_BatchDump", response.model_dump()) # cast-ok: Batch.model_dump returns its typed fields
return LiteLLMBatch(**batch_dict)
def create_batch(
self,
@ -1998,7 +2125,7 @@ class OpenAIBatchesAPI(BaseLLM):
max_retries: Optional[int],
organization: Optional[str],
client: Optional[Union[OpenAI, AsyncOpenAI]] = None,
) -> Union[LiteLLMBatch, Coroutine[Any, Any, LiteLLMBatch]]:
) -> LiteLLMBatch | Coroutine[None, None, LiteLLMBatch]:
openai_client: Optional[Union[OpenAI, AsyncOpenAI]] = self.get_openai_client(
api_key=api_key,
api_base=api_base,
@ -2023,7 +2150,8 @@ class OpenAIBatchesAPI(BaseLLM):
)
response = cast(OpenAI, openai_client).batches.create(**create_batch_data) # type: ignore[arg-type]
return LiteLLMBatch(**response.model_dump())
batch_dict = cast("_BatchDump", response.model_dump()) # cast-ok: Batch.model_dump returns its typed fields
return LiteLLMBatch(**batch_dict)
async def aretrieve_batch(
self,
@ -2032,7 +2160,8 @@ class OpenAIBatchesAPI(BaseLLM):
) -> LiteLLMBatch:
verbose_logger.debug("retrieving batch, args= %s", retrieve_batch_data)
response = await openai_client.batches.retrieve(**retrieve_batch_data) # type: ignore[arg-type]
return LiteLLMBatch(**response.model_dump())
batch_dict = cast("_BatchDump", response.model_dump()) # cast-ok: Batch.model_dump returns its typed fields
return LiteLLMBatch(**batch_dict)
def retrieve_batch(
self,
@ -2068,7 +2197,8 @@ class OpenAIBatchesAPI(BaseLLM):
retrieve_batch_data=retrieve_batch_data, openai_client=openai_client
)
response = cast(OpenAI, openai_client).batches.retrieve(**retrieve_batch_data) # type: ignore[arg-type]
return LiteLLMBatch(**response.model_dump())
batch_dict = cast("_BatchDump", response.model_dump()) # cast-ok: Batch.model_dump returns its typed fields
return LiteLLMBatch(**batch_dict)
async def acancel_batch(
self,
@ -2077,7 +2207,8 @@ class OpenAIBatchesAPI(BaseLLM):
) -> LiteLLMBatch:
verbose_logger.debug("async cancelling batch, args= %s", cancel_batch_data)
response = await openai_client.batches.cancel(**cancel_batch_data)
return LiteLLMBatch(**response.model_dump())
batch_dict = cast("_BatchDump", response.model_dump()) # cast-ok: Batch.model_dump returns its typed fields
return LiteLLMBatch(**batch_dict)
def cancel_batch(
self,
@ -2117,7 +2248,8 @@ class OpenAIBatchesAPI(BaseLLM):
if not isinstance(openai_client, OpenAI):
raise ValueError("OpenAI client is not an instance of OpenAI. Make sure you passed a sync OpenAI client.")
response = openai_client.batches.cancel(**cancel_batch_data)
return LiteLLMBatch(**response.model_dump())
batch_dict = cast("_BatchDump", response.model_dump()) # cast-ok: Batch.model_dump returns its typed fields
return LiteLLMBatch(**batch_dict)
async def alist_batches(
self,
@ -2477,9 +2609,11 @@ class OpenAIAssistantsAPI(BaseLLM):
response_obj: Optional[OpenAIMessage] = None
if getattr(thread_message, "status", None) is None:
thread_message.status = "completed"
response_obj = OpenAIMessage(**thread_message.dict())
message_dict = cast("_MessageDump", thread_message.dict()) # cast-ok: Message.dict returns its typed fields
response_obj = OpenAIMessage(**message_dict)
else:
response_obj = OpenAIMessage(**thread_message.dict())
message_dict = cast("_MessageDump", thread_message.dict()) # cast-ok: Message.dict returns its typed fields
response_obj = OpenAIMessage(**message_dict)
return response_obj
# fmt: off
@ -2556,9 +2690,11 @@ class OpenAIAssistantsAPI(BaseLLM):
response_obj: Optional[OpenAIMessage] = None
if getattr(thread_message, "status", None) is None:
thread_message.status = "completed"
response_obj = OpenAIMessage(**thread_message.dict())
message_dict = cast("_MessageDump", thread_message.dict()) # cast-ok: Message.dict returns its typed fields
response_obj = OpenAIMessage(**message_dict)
else:
response_obj = OpenAIMessage(**thread_message.dict())
message_dict = cast("_MessageDump", thread_message.dict()) # cast-ok: Message.dict returns its typed fields
response_obj = OpenAIMessage(**message_dict)
return response_obj
async def async_get_messages(
@ -2680,7 +2816,10 @@ class OpenAIAssistantsAPI(BaseLLM):
message_thread = await openai_client.beta.threads.create(**data) # type: ignore
return Thread(**message_thread.dict())
thread_dict = cast( # cast-ok: thread dump carries litellm Thread's typed fields
"_ThreadDump", message_thread.dict()
)
return Thread(**thread_dict)
# fmt: off
@ -2766,7 +2905,10 @@ class OpenAIAssistantsAPI(BaseLLM):
message_thread = openai_client.beta.threads.create(**data) # type: ignore
return Thread(**message_thread.dict())
thread_dict = cast( # cast-ok: thread dump carries litellm Thread's typed fields
"_ThreadDump", message_thread.dict()
)
return Thread(**thread_dict)
async def async_get_thread(
self,
@ -2789,7 +2931,8 @@ class OpenAIAssistantsAPI(BaseLLM):
response = await openai_client.beta.threads.retrieve(thread_id=thread_id)
return Thread(**response.dict())
thread_dict = cast("_ThreadDump", response.dict()) # cast-ok: thread dump carries litellm Thread's typed fields
return Thread(**thread_dict)
# fmt: off
@ -2855,7 +2998,8 @@ class OpenAIAssistantsAPI(BaseLLM):
response = openai_client.beta.threads.retrieve(thread_id=thread_id)
return Thread(**response.dict())
thread_dict = cast("_ThreadDump", response.dict()) # cast-ok: thread dump carries litellm Thread's typed fields
return Thread(**thread_dict)
def delete_thread(self):
pass
@ -2912,7 +3056,7 @@ class OpenAIAssistantsAPI(BaseLLM):
tools: Optional[Iterable[AssistantToolParam]],
event_handler: Optional[AssistantEventHandler],
) -> AsyncAssistantStreamManager[AsyncAssistantEventHandler]:
data: Dict[str, Any] = {
data: _AsyncRunThreadStreamData = {
"thread_id": thread_id,
"assistant_id": assistant_id,
"additional_instructions": additional_instructions,
@ -2922,7 +3066,9 @@ class OpenAIAssistantsAPI(BaseLLM):
"tools": tools,
}
if event_handler is not None:
data["event_handler"] = event_handler
data["event_handler"] = cast( # cast-ok: async runs stream requires an async event handler
"AsyncAssistantEventHandler", event_handler
)
return client.beta.threads.runs.stream(**data) # type: ignore
def run_thread_stream(
@ -2937,7 +3083,7 @@ class OpenAIAssistantsAPI(BaseLLM):
tools: Optional[Iterable[AssistantToolParam]],
event_handler: Optional[AssistantEventHandler],
) -> AssistantStreamManager[AssistantEventHandler]:
data: Dict[str, Any] = {
data: _RunThreadStreamData = {
"thread_id": thread_id,
"assistant_id": assistant_id,
"additional_instructions": additional_instructions,

View file

@ -3,7 +3,24 @@ import binascii
import hashlib
import json
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Any, Awaitable, Callable, Dict, Iterable, List, Optional, Set, Union, cast
from typing import (
TYPE_CHECKING,
Any,
Awaitable,
Callable,
Dict,
Iterable,
List,
Mapping,
Optional,
Protocol,
Sequence,
Set,
TypedDict,
TypeVar,
Union,
cast,
)
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
@ -16,7 +33,6 @@ from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
from litellm.proxy._types import (
LiteLLM_MCPServerTable,
LiteLLM_ObjectPermissionTable,
LiteLLM_TeamTable,
MCPApprovalStatus,
MCPEnvVarScope,
MCPSubmissionsSummary,
@ -45,9 +61,19 @@ from litellm.types.llms.custom_http import httpxSpecialProvider
from litellm.types.mcp import MCPCredentials
if TYPE_CHECKING:
from prisma import Prisma
from prisma.models import LiteLLM_MCPServerOAuthClient as PrismaMCPServerOAuthClientRow
from prisma.models import LiteLLM_MCPServerTable as PrismaMCPServerTableRow
from prisma.models import LiteLLM_MCPUserCredentials as PrismaMCPUserCredentialsRow
from prisma.models import LiteLLM_MCPUserEnvVars as PrismaMCPUserEnvVarsRow
from prisma.models import LiteLLM_ObjectPermissionTable as PrismaObjectPermissionTableRow
from prisma.models import LiteLLM_TeamTable as PrismaTeamTableRow
from prisma.models import LiteLLM_VerificationToken as PrismaVerificationTokenRow
from litellm.models.mcp_server import MCPEnvVar
from litellm.types.mcp_server.mcp_server_manager import MCPServer
_AUTH_FLOW_SCOPED_FIELDS: frozenset = frozenset(
_AUTH_FLOW_SCOPED_FIELDS: frozenset[str] = frozenset(
{
"issuer",
"authorization_url",
@ -76,7 +102,7 @@ def _blank_to_none(value: Optional[str]) -> Optional[str]:
# the current code has never written — a cleared column can then never be
# silently resurrected by a stale blob copy. These keys are stored plaintext
# (endpoints/identifiers, not secrets), so values lift as-is.
_TOKEN_EXCHANGE_COLUMN_FIELDS: frozenset = frozenset(
_TOKEN_EXCHANGE_COLUMN_FIELDS: frozenset[str] = frozenset(
{
"token_exchange_endpoint",
"audience",
@ -89,10 +115,122 @@ _TOKEN_EXCHANGE_COLUMN_FIELDS: frozenset = frozenset(
# OAuth app (client_id/client_secret) plus the same authorize relay, and neither mints anything the
# gateway keeps. So a switch WITHIN this class must preserve the stored app, unlike a cross-class
# switch (e.g. an oauth2 row whose client may be DCR-minted and is not reusable elsewhere).
_CLIENT_FORWARDED_AUTH_TYPES: frozenset = frozenset({"true_passthrough", "oauth_delegate"})
_CLIENT_FORWARDED_AUTH_TYPES: frozenset[str] = frozenset({"true_passthrough", "oauth_delegate"})
# Minted token material that must never survive a client rotation on a persisted row.
_MINTED_TOKEN_CREDENTIAL_FIELDS: frozenset = frozenset({"access_token", "refresh_token", "expires_in"})
_MINTED_TOKEN_CREDENTIAL_FIELDS: frozenset[str] = frozenset({"access_token", "refresh_token", "expires_in"})
_RowT_co = TypeVar("_RowT_co", covariant=True)
class _PrismaTable(Protocol[_RowT_co]):
async def find_unique(
self,
where: Mapping[str, object],
include: Mapping[str, bool] | None = None,
) -> _RowT_co | None: ...
async def find_many(
self,
where: Mapping[str, object] | None = None,
include: Mapping[str, bool] | None = None,
order: Mapping[str, str] | None = None,
take: int | None = None,
) -> Sequence[_RowT_co]: ...
async def create(self, data: Mapping[str, object]) -> _RowT_co: ...
async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> _RowT_co: ...
async def upsert(self, where: Mapping[str, object], data: Mapping[str, Mapping[str, object]]) -> _RowT_co: ...
async def delete(self, where: Mapping[str, object]) -> _RowT_co | None: ...
async def delete_many(self, where: Mapping[str, object]) -> int: ...
def _mcp_server_table(prisma_client: PrismaClient) -> "_PrismaTable[PrismaMCPServerTableRow]":
return cast( # cast-ok: repository ``table`` is typed Any; narrow it to the prisma table surface used here
"_PrismaTable[PrismaMCPServerTableRow]",
MCPServerRepository(prisma_client).table,
)
def _verification_token_table(prisma_client: PrismaClient) -> "_PrismaTable[PrismaVerificationTokenRow]":
return cast( # cast-ok: repository ``table`` is typed Any; narrow it to the prisma table surface used here
"_PrismaTable[PrismaVerificationTokenRow]",
VerificationTokenRepository(prisma_client).table,
)
def _team_table(prisma_client: PrismaClient) -> "_PrismaTable[PrismaTeamTableRow]":
return cast( # cast-ok: repository ``table`` is typed Any; narrow it to the prisma table surface used here
"_PrismaTable[PrismaTeamTableRow]",
TeamRepository(prisma_client).table,
)
def _object_permission_table(prisma_client: PrismaClient) -> "_PrismaTable[PrismaObjectPermissionTableRow]":
return cast( # cast-ok: repository ``table`` is typed Any; narrow it to the prisma table surface used here
"_PrismaTable[PrismaObjectPermissionTableRow]",
ObjectPermissionRepository(prisma_client).table,
)
def _mcp_user_credentials_table(prisma_client: PrismaClient) -> "_PrismaTable[PrismaMCPUserCredentialsRow]":
return cast( # cast-ok: repository ``table`` is typed Any; narrow it to the prisma table surface used here
"_PrismaTable[PrismaMCPUserCredentialsRow]",
MCPUserCredentialsRepository(prisma_client).table,
)
def _mcp_oauth_client_table(prisma_client: PrismaClient) -> "_PrismaTable[PrismaMCPServerOAuthClientRow]":
return cast( # cast-ok: repository ``table`` is typed Any; narrow it to the prisma table surface used here
"_PrismaTable[PrismaMCPServerOAuthClientRow]",
MCPServerOAuthClientRepository(prisma_client).table,
)
def _prisma_db(prisma_client: PrismaClient) -> "Prisma":
return cast( # cast-ok: PrismaClient.db is an untyped proxy around the generated prisma client
"Prisma", prisma_client.db
)
def _mcp_user_env_vars_table(prisma_client: PrismaClient) -> "_PrismaTable[PrismaMCPUserEnvVarsRow]":
return cast( # cast-ok: the generated table client methods reference types pyright cannot load; narrow to the surface used here
"_PrismaTable[PrismaMCPUserEnvVarsRow]",
_prisma_db(prisma_client).litellm_mcpuserenvvars,
)
def _credentials_blob(value: object) -> str | Mapping[str, object] | None:
return cast( # cast-ok: prisma types Json columns opaquely; at runtime the column holds the deserialized blob
"str | Mapping[str, object] | None", value
)
def _env_vars_blob(value: object) -> Sequence[Mapping[str, str]] | None:
return cast( # cast-ok: prisma types Json columns opaquely; at runtime the column holds the deserialized blob
"Sequence[Mapping[str, str]] | None", value
)
def _loads_json_object(value: str) -> object:
return cast( # cast-ok: json.loads returns Any; the parsed payload is handled as an opaque object
object, json.loads(value)
)
def _parse_credentials_blob_dict(
blob: str | Mapping[str, object],
) -> dict[str, object]: # mutable-ok: callers pop lifted token-exchange keys off the returned copy
if isinstance(blob, str):
return cast( # cast-ok: the credentials column always stores a serialized JSON object
dict[str, object], json.loads(blob)
)
return dict(blob)
def _credential_auth_class(auth_type: Optional[str]) -> Optional[str]:
@ -104,7 +242,7 @@ def _credential_auth_class(auth_type: Optional[str]) -> Optional[str]:
return auth_type
def _drop_stale_minted_on_client_rotation(merged: Dict[str, Any], new_creds: Dict[str, Any]) -> Dict[str, Any]:
def _drop_stale_minted_on_client_rotation(merged: dict[str, object], new_creds: dict[str, object]) -> dict[str, object]:
"""When the update rotates the client, drop stale minted token keys it did not itself set, so an old
app's access/refresh token never rides forward under the new client. A no-op when no client key changed."""
if "client_id" not in new_creds and "client_secret" not in new_creds:
@ -114,13 +252,13 @@ def _drop_stale_minted_on_client_rotation(merged: Dict[str, Any], new_creds: Dic
}
def _is_global_env_var_scope(scope: Any) -> bool:
def _is_global_env_var_scope(scope: object) -> bool:
"""``scope="user"`` entries are placeholders the user fills in; everything
else (including a missing scope) is an admin-supplied global value."""
return scope != MCPEnvVarScope.user and scope != "user"
def _encrypt_global_env_var_values(env_vars: Iterable[Dict[str, Any]]) -> None:
def _encrypt_global_env_var_values(env_vars: Iterable[dict[str, str]]) -> None:
"""Encrypt ``scope="global"`` env var values in place before persisting.
Global values hold admin-supplied secrets (API keys, passwords) that get
@ -136,7 +274,11 @@ def _encrypt_global_env_var_values(env_vars: Iterable[Dict[str, Any]]) -> None:
entry["value"] = encrypt_value_helper(value)
def decrypt_global_env_var_values(env_vars: Optional[Iterable[Any]]) -> None:
def decrypt_global_env_var_values(
env_vars: Optional[
Iterable[Union[dict[str, str], "MCPEnvVar"]]
], # mutable-ok: each entry's value is decrypted in place
) -> None:
"""Decrypt ``scope="global"`` env var values in place after reading the DB.
Accepts ``MCPEnvVar`` models (``LiteLLM_MCPServerTable``) or plain dicts
@ -175,7 +317,7 @@ def decrypt_global_env_var_values(env_vars: Optional[Iterable[Any]]) -> None:
entry.value = decrypted
def _decrypt_env_vars_on_returned_row(row: Any) -> None:
def _decrypt_env_vars_on_returned_row(row: object) -> None:
"""Decrypt ``scope="global"`` env var values on a row returned by Prisma create/update.
Prisma may hand back ``env_vars`` either as a parsed list (the common case for
@ -187,16 +329,21 @@ def _decrypt_env_vars_on_returned_row(row: Any) -> None:
write the decrypted list back onto the row so downstream consumers see plain
values.
"""
env_vars = getattr(row, "env_vars", None)
env_vars = cast( # cast-ok: prisma returns the Json column as a parsed list or a raw JSON string
"str | list[dict[str, str]] | None", getattr(row, "env_vars", None)
)
if env_vars is None:
return
if isinstance(env_vars, str):
try:
env_vars = json.loads(env_vars)
parsed = _loads_json_object(env_vars)
except (json.JSONDecodeError, TypeError):
return
if not isinstance(env_vars, list):
if not isinstance(parsed, list):
return
env_vars = cast( # cast-ok: the parsed payload is the stored env var entry list
"list[dict[str, str]]", parsed
)
try:
setattr(row, "env_vars", env_vars)
except (AttributeError, TypeError):
@ -205,8 +352,8 @@ def _decrypt_env_vars_on_returned_row(row: Any) -> None:
def _reencrypt_global_env_var_values(
env_vars: Optional[Iterable[Any]], new_encryption_key: str
) -> Optional[List[Dict[str, Any]]]:
env_vars: Iterable[Mapping[str, str]] | None, new_encryption_key: str
) -> Sequence[Mapping[str, str]] | None:
"""Re-encrypt ``scope="global"`` env var values for master-key rotation.
Each global value is decrypted with the current salt key and re-encrypted
@ -219,7 +366,9 @@ def _reencrypt_global_env_var_values(
return None
if isinstance(env_vars, str):
try:
env_vars = json.loads(env_vars)
env_vars = cast( # cast-ok: the env var column stores a serialized list of entries
"list[dict[str, str]]", json.loads(env_vars)
)
except (json.JSONDecodeError, TypeError):
return None
if not env_vars:
@ -438,12 +587,12 @@ async def get_all_mcp_servers(
Pass approval_status=None to return all servers regardless of approval state.
"""
try:
where: Dict[str, Any] = {}
where: dict[str, str] = {}
if approval_status is not None:
where["approval_status"] = approval_status
mcp_servers = await MCPServerRepository(prisma_client).table.find_many(where=where if where else {})
mcp_servers = await _mcp_server_table(prisma_client).find_many(where=where if where else {})
tables = [LiteLLM_MCPServerTable(**mcp_server.model_dump()) for mcp_server in mcp_servers]
tables = [LiteLLM_MCPServerTable.model_validate(mcp_server.model_dump()) for mcp_server in mcp_servers]
for table in tables:
decrypt_global_env_var_values(table.env_vars)
return tables
@ -458,14 +607,14 @@ async def get_mcp_server(prisma_client: PrismaClient, server_id: str) -> Optiona
"""
Returns the matching mcp server from the db iff exists
"""
mcp_server: Optional[LiteLLM_MCPServerTable] = await MCPServerRepository(prisma_client).table.find_unique(
mcp_server = await _mcp_server_table(prisma_client).find_unique(
where={
"server_id": server_id,
}
)
if mcp_server is None:
return None
table = LiteLLM_MCPServerTable(**mcp_server.model_dump())
table = LiteLLM_MCPServerTable.model_validate(mcp_server.model_dump())
decrypt_global_env_var_values(table.env_vars)
return table
@ -474,14 +623,14 @@ async def get_mcp_servers(prisma_client: PrismaClient, server_ids: Iterable[str]
"""
Returns the matching mcp servers from the db with the server_ids
"""
_mcp_servers: List[LiteLLM_MCPServerTable] = await MCPServerRepository(prisma_client).table.find_many(
_mcp_servers = await _mcp_server_table(prisma_client).find_many(
where={
"server_id": {"in": server_ids},
"server_id": {"in": list(server_ids)},
}
)
final_mcp_servers: List[LiteLLM_MCPServerTable] = []
for _mcp_server in _mcp_servers:
table = LiteLLM_MCPServerTable(**_mcp_server.model_dump())
table = LiteLLM_MCPServerTable.model_validate(_mcp_server.model_dump())
decrypt_global_env_var_values(table.env_vars)
final_mcp_servers.append(table)
@ -492,7 +641,7 @@ async def get_mcp_servers_by_verificationtoken(prisma_client: PrismaClient, toke
"""
Returns the mcp servers from the db for the verification token
"""
verification_token_record: LiteLLM_TeamTable = await VerificationTokenRepository(prisma_client).table.find_unique(
verification_token_record = await _verification_token_table(prisma_client).find_unique(
where={
"token": token,
},
@ -511,7 +660,7 @@ async def get_mcp_servers_by_team(prisma_client: PrismaClient, team_id: str) ->
"""
Returns the mcp servers from the db for the team id
"""
team_record: LiteLLM_TeamTable = await TeamRepository(prisma_client).table.find_unique(
team_record = await _team_table(prisma_client).find_unique(
where={
"team_id": team_id,
},
@ -561,7 +710,7 @@ async def get_objectpermissions_for_mcp_server(
"""
Get all the object permissions records and the associated team and verficiationtoken records that have access to the mcp server
"""
object_permission_records = await ObjectPermissionRepository(prisma_client).table.find_many(
object_permission_records = await _object_permission_table(prisma_client).find_many(
where={
"mcp_servers": {"has": mcp_server_id},
},
@ -571,21 +720,23 @@ async def get_objectpermissions_for_mcp_server(
},
)
return object_permission_records
return cast( # cast-ok: callers consume the prisma rows under the table-model annotation; runtime is unchanged
list[LiteLLM_ObjectPermissionTable], object_permission_records
)
async def get_virtualkeys_for_mcp_server(prisma_client: PrismaClient, server_id: str) -> List:
async def get_virtualkeys_for_mcp_server(
prisma_client: PrismaClient, server_id: str
) -> Sequence["PrismaVerificationTokenRow"]:
"""
Get all the virtual keys that have access to the mcp server
"""
virtual_keys = await VerificationTokenRepository(prisma_client).table.find_many(
virtual_keys = await _verification_token_table(prisma_client).find_many(
where={
"mcp_servers": {"has": server_id},
},
)
if virtual_keys is None:
return []
return virtual_keys
@ -626,7 +777,7 @@ async def delete_mcp_server(
Returns the deleted mcp server record if it exists, otherwise None
"""
deleted_server = await MCPServerRepository(prisma_client).table.delete(
deleted_server = await _mcp_server_table(prisma_client).delete(
where={
"server_id": server_id,
},
@ -634,9 +785,7 @@ async def delete_mcp_server(
if deleted_server is not None:
credential_user_ids: List[str] = []
try:
credential_rows = await prisma_client.db.litellm_mcpusercredentials.find_many(
where={"server_id": server_id}
)
credential_rows = await _mcp_user_credentials_table(prisma_client).find_many(where={"server_id": server_id})
credential_user_ids = [row.user_id for row in credential_rows]
except Exception as e: # noqa: BLE001 - enumeration is best-effort; cached tokens expire by TTL
verbose_proxy_logger.warning(
@ -645,9 +794,9 @@ async def delete_mcp_server(
e,
)
for model, label in (
(prisma_client.db.litellm_mcpusercredentials, "credential"),
(prisma_client.db.litellm_mcpuserenvvars, "env var"),
(prisma_client.db.litellm_mcpserveroauthclient, "OAuth client"),
(_mcp_user_credentials_table(prisma_client), "credential"),
(_mcp_user_env_vars_table(prisma_client), "env var"),
(_mcp_oauth_client_table(prisma_client), "OAuth client"),
):
try:
await model.delete_many(where={"server_id": server_id})
@ -668,7 +817,9 @@ async def delete_mcp_server(
invalidate_token_cache = global_mcp_server_manager.invalidate_user_oauth_token_cache
for user_id in credential_user_ids:
await invalidate_token_cache(user_id, server_id)
return deleted_server
return cast( # cast-ok: callers consume the prisma row under the table-model annotation; runtime is unchanged
LiteLLM_MCPServerTable | None, deleted_server
)
async def create_mcp_server(
@ -687,12 +838,12 @@ async def create_mcp_server(
data_dict["created_by"] = touched_by
data_dict["updated_by"] = touched_by
new_mcp_server = await MCPServerRepository(prisma_client).table.create(
data=data_dict # type: ignore
)
new_mcp_server = await _mcp_server_table(prisma_client).create(data=data_dict)
_decrypt_env_vars_on_returned_row(new_mcp_server)
return new_mcp_server
return cast( # cast-ok: callers consume the prisma row under the table-model annotation; runtime is unchanged
LiteLLM_MCPServerTable, new_mcp_server
)
async def update_mcp_server(
@ -704,8 +855,6 @@ async def update_mcp_server(
"""
Update a new mcp server record in the db
"""
import json
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
# Use helper to prepare data with proper JSON serialization.
@ -723,7 +872,8 @@ async def update_mcp_server(
url_provided = "url" in data_dict and data_dict["url"] is not None
issuer_provided = "issuer" in data_dict
if data.auth_type or has_credentials or explicit_te_write or url_provided or issuer_provided:
existing = await MCPServerRepository(prisma_client).table.find_unique(where={"server_id": data.server_id})
existing = await _mcp_server_table(prisma_client).find_unique(where={"server_id": data.server_id})
existing_credentials = _credentials_blob(existing.credentials) if existing is not None else None
auth_type_changed = bool(
data.auth_type
@ -762,10 +912,8 @@ async def update_mcp_server(
# place, and the next credentials update's migrate-on-write would silently
# repopulate the column the admin just cleared. (When credentials ARE in the
# update, the merge below performs the same migration.)
if explicit_te_write and "credentials" not in data_dict and existing is not None and existing.credentials:
existing_creds = (
json.loads(existing.credentials) if isinstance(existing.credentials, str) else dict(existing.credentials)
)
if explicit_te_write and "credentials" not in data_dict and existing is not None and existing_credentials:
existing_creds = _parse_credentials_blob_dict(existing_credentials)
if _TOKEN_EXCHANGE_COLUMN_FIELDS & existing_creds.keys():
for te_field in _TOKEN_EXCHANGE_COLUMN_FIELDS:
legacy_value = existing_creds.pop(te_field, None)
@ -777,23 +925,15 @@ async def update_mcp_server(
# Without this, a partial credential update (e.g. changing only region)
# would wipe encrypted secrets that the UI cannot display back.
if "credentials" in data_dict and data_dict["credentials"] is not None:
if existing and existing.credentials:
if existing and existing_credentials:
# Only merge when the credential CLASS is unchanged. A cross-class switch
# (e.g. oauth2 → api_key, or oauth2 → true_passthrough) replaces credentials
# entirely to avoid stale secrets from the previous class lingering; a switch
# within the client-forwarded class (true_passthrough ↔ oauth_delegate) keeps
# the same declared app and so must merge, not replace.
if not auth_type_changed:
existing_creds = (
json.loads(existing.credentials)
if isinstance(existing.credentials, str)
else dict(existing.credentials)
)
new_creds = (
json.loads(data_dict["credentials"])
if isinstance(data_dict["credentials"], str)
else dict(data_dict["credentials"])
)
existing_creds = _parse_credentials_blob_dict(existing_credentials)
new_creds = _parse_credentials_blob_dict(data_dict["credentials"])
# New values override existing; existing keys not in update are preserved. A client
# rotation additionally drops the previous app's stale minted token keys.
merged = _drop_stale_minted_on_client_rotation({**existing_creds, **new_creds}, new_creds)
@ -823,13 +963,15 @@ async def update_mcp_server(
data_dict["credentials"] = Json(None)
updated_mcp_server = await MCPServerRepository(prisma_client).table.update(
updated_mcp_server = await _mcp_server_table(prisma_client).update(
where={"server_id": data.server_id},
data=data_dict, # type: ignore
data=data_dict,
)
_decrypt_env_vars_on_returned_row(updated_mcp_server)
return updated_mcp_server
return cast( # cast-ok: callers consume the prisma row under the table-model annotation; runtime is unchanged
LiteLLM_MCPServerTable, updated_mcp_server
)
async def get_mcp_server_oauth_client_credentials(prisma_client: PrismaClient, server_id: str) -> object | None:
@ -838,7 +980,7 @@ async def get_mcp_server_oauth_client_credentials(prisma_client: PrismaClient, s
LiteLLM_MCPServerTable row, so their dynamically registered client lives here keyed
by server_id. The returned value is the raw credentials blob for
``_get_persisted_dcr_credentials`` to parse."""
row = await MCPServerOAuthClientRepository(prisma_client).table.find_unique(where={"server_id": server_id})
row = await _mcp_oauth_client_table(prisma_client).find_unique(where={"server_id": server_id})
if row is None:
return None
return row.credentials
@ -856,7 +998,7 @@ async def upsert_mcp_server_oauth_client_credentials(
encrypted = encrypt_credentials(credentials=dict(credentials), encryption_key=_get_salt_key())
blob = safe_dumps(encrypted)
await MCPServerOAuthClientRepository(prisma_client).table.upsert(
await _mcp_oauth_client_table(prisma_client).upsert(
where={"server_id": server_id},
data={
"create": {"server_id": server_id, "credentials": blob},
@ -883,17 +1025,17 @@ def _reencrypt_mcp_credentials_blob(credentials: object, new_master_key: str) ->
async def rotate_mcp_server_credentials_master_key(prisma_client: PrismaClient, touched_by: str, new_master_key: str):
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps # noqa: PLC0415 # avoids circular import
mcp_servers = await MCPServerRepository(prisma_client).table.find_many()
mcp_servers = await _mcp_server_table(prisma_client).find_many()
updated = 0
for mcp_server in mcp_servers:
update_data: Dict[str, Any] = {}
update_data: dict[str, str] = {}
rotated_credentials = _reencrypt_mcp_credentials_blob(mcp_server.credentials, new_master_key)
if rotated_credentials is not None:
update_data["credentials"] = rotated_credentials
rotated_env_vars = _reencrypt_global_env_var_values(mcp_server.env_vars, new_master_key)
rotated_env_vars = _reencrypt_global_env_var_values(_env_vars_blob(mcp_server.env_vars), new_master_key)
if rotated_env_vars is not None:
update_data["env_vars"] = safe_dumps(rotated_env_vars)
@ -901,19 +1043,19 @@ async def rotate_mcp_server_credentials_master_key(prisma_client: PrismaClient,
continue
update_data["updated_by"] = touched_by
await MCPServerRepository(prisma_client).table.update(
await _mcp_server_table(prisma_client).update(
where={"server_id": mcp_server.server_id},
data=update_data,
)
updated += 1
oauth_clients = await MCPServerOAuthClientRepository(prisma_client).table.find_many()
oauth_clients = await _mcp_oauth_client_table(prisma_client).find_many()
oauth_updated = 0
for oauth_client in oauth_clients:
rotated_credentials = _reencrypt_mcp_credentials_blob(oauth_client.credentials, new_master_key)
if rotated_credentials is None:
continue
await MCPServerOAuthClientRepository(prisma_client).table.update(
await _mcp_oauth_client_table(prisma_client).update(
where={"server_id": oauth_client.server_id},
data={"credentials": rotated_credentials},
)
@ -959,7 +1101,7 @@ def _decode_oauth_payload(stored: str) -> Optional[Dict[str, Any]]:
if decoded is None:
return None
try:
parsed = json.loads(decoded)
parsed = _loads_json_object(decoded)
except (ValueError, TypeError):
return None
if isinstance(parsed, dict) and parsed.get("type") == "oauth2":
@ -975,7 +1117,7 @@ async def rotate_mcp_user_credentials_master_key(prisma_client: PrismaClient, ne
under the new master key. Rows that are unreadable under both paths
are logged and skipped so one corrupt row does not abort the rotation.
"""
rows = await MCPUserCredentialsRepository(prisma_client).table.find_many()
rows = await _mcp_user_credentials_table(prisma_client).find_many()
rotated = 0
skipped = 0
for row in rows:
@ -990,7 +1132,7 @@ async def rotate_mcp_user_credentials_master_key(prisma_client: PrismaClient, ne
skipped += 1
continue
re_encrypted = encrypt_value_helper(plaintext, new_encryption_key=new_master_key)
await MCPUserCredentialsRepository(prisma_client).table.update(
await _mcp_user_credentials_table(prisma_client).update(
where={
"user_id_server_id": {
"user_id": row.user_id,
@ -1015,7 +1157,7 @@ async def rotate_mcp_user_env_vars_master_key(prisma_client: PrismaClient, new_m
skipped so one corrupt row does not abort the rotation nor overwrite values
that may still be recoverable.
"""
rows = await prisma_client.db.litellm_mcpuserenvvars.find_many()
rows = await _mcp_user_env_vars_table(prisma_client).find_many()
rotated = 0
skipped = 0
for row in rows:
@ -1034,7 +1176,7 @@ async def rotate_mcp_user_env_vars_master_key(prisma_client: PrismaClient, new_m
skipped += 1
continue
re_encrypted = encrypt_value_helper(plaintext, new_encryption_key=new_master_key)
await prisma_client.db.litellm_mcpuserenvvars.update(
await _mcp_user_env_vars_table(prisma_client).update(
where={
"user_id_server_id": {
"user_id": row.user_id,
@ -1060,7 +1202,7 @@ async def store_user_credential(
"""Store a user credential for a BYOK MCP server."""
encoded = encrypt_value_helper(credential)
await MCPUserCredentialsRepository(prisma_client).table.upsert(
await _mcp_user_credentials_table(prisma_client).upsert(
where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}},
data={
"create": {
@ -1080,7 +1222,7 @@ async def get_user_credential(
) -> Optional[str]:
"""Return credential for a user+server pair, or None."""
row = await MCPUserCredentialsRepository(prisma_client).table.find_unique(
row = await _mcp_user_credentials_table(prisma_client).find_unique(
where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}}
)
if row is None:
@ -1094,7 +1236,7 @@ async def has_user_credential(
server_id: str,
) -> bool:
"""Return True if the user has a stored credential for this server."""
row = await MCPUserCredentialsRepository(prisma_client).table.find_unique(
row = await _mcp_user_credentials_table(prisma_client).find_unique(
where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}}
)
return row is not None
@ -1106,7 +1248,7 @@ async def delete_user_credential(
server_id: str,
) -> None:
"""Delete the user's stored credential for a BYOK MCP server."""
await MCPUserCredentialsRepository(prisma_client).table.delete(
await _mcp_user_credentials_table(prisma_client).delete(
where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}}
)
@ -1135,7 +1277,7 @@ async def store_user_oauth_credential(
if expires_in is not None:
expires_at = (datetime.now(timezone.utc) + timedelta(seconds=expires_in)).isoformat()
payload: Dict[str, Any] = {
payload: dict[str, object] = {
"type": "oauth2",
"access_token": access_token,
"connected_at": datetime.now(timezone.utc).isoformat(),
@ -1151,7 +1293,7 @@ async def store_user_oauth_credential(
# Skip the guard when the caller knows the row is already an OAuth2 credential
# (e.g. during token refresh), saving an extra DB round-trip.
if not skip_byok_guard:
existing = await MCPUserCredentialsRepository(prisma_client).table.find_unique(
existing = await _mcp_user_credentials_table(prisma_client).find_unique(
where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}}
)
if existing is not None and _decode_oauth_payload(existing.credential_b64) is None:
@ -1166,7 +1308,7 @@ async def store_user_oauth_credential(
)
encoded = encrypt_value_helper(json.dumps(payload))
await MCPUserCredentialsRepository(prisma_client).table.upsert(
await _mcp_user_credentials_table(prisma_client).upsert(
where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}},
data={
"create": {
@ -1207,7 +1349,7 @@ async def get_user_oauth_credential(
) -> Optional[Dict[str, Any]]:
"""Return the decoded OAuth2 payload dict for a user+server pair, or None."""
row = await MCPUserCredentialsRepository(prisma_client).table.find_unique(
row = await _mcp_user_credentials_table(prisma_client).find_unique(
where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}}
)
if row is None:
@ -1221,7 +1363,7 @@ async def list_user_oauth_credentials(
) -> List[Dict[str, Any]]:
"""Return all OAuth2 credential payloads for a user, tagged with server_id."""
rows = await MCPUserCredentialsRepository(prisma_client).table.find_many(where={"user_id": user_id})
rows = await _mcp_user_credentials_table(prisma_client).find_many(where={"user_id": user_id})
results: List[Dict[str, Any]] = []
for row in rows:
payload = _decode_oauth_payload(row.credential_b64)
@ -1260,7 +1402,7 @@ def mcp_oauth_token_identity(server: object) -> tuple[object, ...]:
creds = getattr(server, "credentials", None)
if isinstance(creds, str):
try:
parsed: object = json.loads(creds)
parsed: object = _loads_json_object(creds)
except ValueError:
parsed = None
else:
@ -1302,12 +1444,12 @@ async def purge_user_oauth_credentials_for_server(
invalidate_token_cache is injectable for tests; it defaults to the manager's shared
invalidate_user_oauth_token_cache, the single invalidation point for per-user tokens."""
repo = MCPUserCredentialsRepository(prisma_client)
rows = await repo.table.find_many(where={"server_id": server_id})
repo = _mcp_user_credentials_table(prisma_client)
rows = await repo.find_many(where={"server_id": server_id})
oauth_rows = [row for row in rows if _decode_oauth_payload(row.credential_b64) is not None]
if not oauth_rows:
return 0
deleted_count = await repo.table.delete_many(
deleted_count = await repo.delete_many(
where={"server_id": server_id, "user_id": {"in": [row.user_id for row in oauth_rows]}}
)
if invalidate_token_cache is None:
@ -1330,10 +1472,17 @@ async def purge_user_oauth_credentials_for_server(
return deleted_count
class _OAuthTokenResponse(TypedDict, total=False):
access_token: str
refresh_token: str
expires_in: int
scope: str
async def refresh_user_oauth_token(
prisma_client: PrismaClient,
user_id: str,
server: Any,
server: "MCPServer",
cred: Dict[str, Any],
) -> Optional[Dict[str, Any]]:
"""Attempt to refresh a per-user OAuth2 token using its stored refresh_token.
@ -1384,7 +1533,9 @@ async def refresh_user_oauth_token(
data=token_data,
)
response.raise_for_status()
body: Dict[str, Any] = response.json()
body = cast( # cast-ok: the token endpoint returns the RFC 6749 token-response JSON object
"_OAuthTokenResponse", response.json()
)
except Exception as exc:
verbose_proxy_logger.warning(
"refresh_user_oauth_token: refresh request failed for user=%s server=%s: %s",
@ -1439,7 +1590,7 @@ async def refresh_user_oauth_token(
async def resolve_valid_user_oauth_token(
user_id: str,
server: Any,
server: "MCPServer",
cred: Optional[Dict[str, Any]],
prisma_client: Optional[PrismaClient] = None,
) -> Optional[Dict[str, Any]]:
@ -1568,7 +1719,7 @@ async def get_active_submitted_mcp_server_ids_for_user(
if not user_id:
return []
rows = await MCPServerRepository(prisma_client).table.find_many(
rows = await _mcp_server_table(prisma_client).find_many(
where={
"submitted_by": user_id,
"approval_status": MCPApprovalStatus.active,
@ -1584,7 +1735,7 @@ async def approve_mcp_server(
) -> LiteLLM_MCPServerTable:
"""Set approval_status=active and record reviewed_at."""
now = datetime.now(timezone.utc)
updated = await MCPServerRepository(prisma_client).table.update(
updated = await _mcp_server_table(prisma_client).update(
where={"server_id": server_id},
data={
"approval_status": MCPApprovalStatus.active,
@ -1592,7 +1743,7 @@ async def approve_mcp_server(
"updated_by": touched_by,
},
)
table = LiteLLM_MCPServerTable(**updated.model_dump())
table = LiteLLM_MCPServerTable.model_validate(updated.model_dump())
decrypt_global_env_var_values(table.env_vars)
return table
@ -1605,18 +1756,18 @@ async def reject_mcp_server(
) -> LiteLLM_MCPServerTable:
"""Set approval_status=rejected, record reviewed_at and review_notes."""
now = datetime.now(timezone.utc)
data: Dict[str, Any] = {
data: dict[str, object] = {
"approval_status": MCPApprovalStatus.rejected,
"reviewed_at": now,
"updated_by": touched_by,
}
if review_notes is not None:
data["review_notes"] = review_notes
updated = await MCPServerRepository(prisma_client).table.update(
updated = await _mcp_server_table(prisma_client).update(
where={"server_id": server_id},
data=data,
)
table = LiteLLM_MCPServerTable(**updated.model_dump())
table = LiteLLM_MCPServerTable.model_validate(updated.model_dump())
decrypt_global_env_var_values(table.env_vars)
return table
@ -1629,12 +1780,12 @@ async def get_mcp_submissions(
along with a summary count breakdown by approval_status.
Mirrors get_guardrail_submissions() from guardrail_endpoints.py.
"""
rows = await MCPServerRepository(prisma_client).table.find_many(
rows = await _mcp_server_table(prisma_client).find_many(
where={"submitted_at": {"not": None}},
order={"submitted_at": "desc"},
take=500, # safety cap; paginate if needed in a future iteration
)
items = [LiteLLM_MCPServerTable(**r.model_dump()) for r in rows]
items = [LiteLLM_MCPServerTable.model_validate(r.model_dump()) for r in rows]
for item in items:
decrypt_global_env_var_values(item.env_vars)
@ -1671,7 +1822,7 @@ def _decode_user_env_vars(stored: str) -> Dict[str, str]:
)
return {}
try:
parsed = json.loads(decrypted)
parsed = _loads_json_object(decrypted)
except (ValueError, TypeError):
return {}
if not isinstance(parsed, dict):
@ -1685,7 +1836,7 @@ async def get_user_env_vars(
server_id: str,
) -> Dict[str, str]:
"""Return the calling user's env var dict for ``server_id`` (empty if none)."""
row = await prisma_client.db.litellm_mcpuserenvvars.find_unique(
row = await _mcp_user_env_vars_table(prisma_client).find_unique(
where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}}
)
if row is None:
@ -1705,7 +1856,7 @@ async def get_user_env_vars_bulk(
ids = list(server_ids)
if not ids:
return {}
rows = await prisma_client.db.litellm_mcpuserenvvars.find_many(where={"user_id": user_id, "server_id": {"in": ids}})
rows = await _mcp_user_env_vars_table(prisma_client).find_many(where={"user_id": user_id, "server_id": {"in": ids}})
return {row.server_id: _decode_user_env_vars(row.values_b64) for row in rows}
@ -1730,15 +1881,18 @@ async def merge_user_env_vars(
"big",
signed=True,
)
async with prisma_client.db.tx() as tx:
async with _prisma_db(prisma_client).tx() as tx:
await tx.execute_raw("SELECT pg_advisory_xact_lock($1::bigint)", lock_key)
row = await tx.litellm_mcpuserenvvars.find_unique(
tx_env_vars_table = cast( # cast-ok: the generated table client methods reference types pyright cannot load; narrow to the surface used here
"_PrismaTable[PrismaMCPUserEnvVarsRow]", tx.litellm_mcpuserenvvars
)
row = await tx_env_vars_table.find_unique(
where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}}
)
existing = _decode_user_env_vars(row.values_b64) if row is not None else {}
merged = {k: v for k, v in {**existing, **updates}.items() if k in allowed}
encoded = encrypt_value_helper(json.dumps(merged))
await tx.litellm_mcpuserenvvars.upsert(
await tx_env_vars_table.upsert(
where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}},
data={
"create": {
@ -1762,4 +1916,4 @@ async def delete_user_env_vars(
Uses ``delete_many`` so a missing row is a no-op; real DB errors still
propagate to the caller instead of being silently swallowed.
"""
await prisma_client.db.litellm_mcpuserenvvars.delete_many(where={"user_id": user_id, "server_id": server_id})
await _mcp_user_env_vars_table(prisma_client).delete_many(where={"user_id": user_id, "server_id": server_id})

View file

@ -13,7 +13,24 @@ import asyncio
import math
import re
import time
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Type, Union, cast
from datetime import datetime
from typing import (
TYPE_CHECKING,
Any,
Dict,
Iterable,
List,
Literal,
Mapping,
Optional,
Protocol,
Sequence,
Type,
TypedDict,
TypeVar,
Union,
cast,
)
from fastapi import HTTPException, Request, status
from pydantic import BaseModel
@ -119,6 +136,393 @@ db_cache_expiry = DEFAULT_IN_MEMORY_TTL # refresh every 5s
all_routes = LiteLLMRoutes.openai_routes.value + LiteLLMRoutes.management_routes.value
class _BudgetRowFields(TypedDict):
budget_id: str | None
soft_budget: float | None
max_budget: float | None
max_parallel_requests: int | None
tpm_limit: int | None
rpm_limit: int | None
model_max_budget: Optional[
dict[str, object]
] # mutable-ok: splatted into LiteLLM_BudgetTable, which declares this field as a dict
budget_duration: str | None
allowed_models: Optional[
list[str]
] # mutable-ok: splatted into LiteLLM_BudgetTable, which declares this field as a list
class _EndUserRowFields(TypedDict):
user_id: str
blocked: bool
alias: str | None
spend: float
allowed_model_region: Literal["eu", "us"] | None
default_model: str | None
budget_id: str | None
class _TagRowFields(TypedDict):
tag_name: str
description: str | None
models: list[str] # mutable-ok: splatted into LiteLLM_TagTable, which declares this field as a list
model_info: Optional[
dict[str, object]
] # mutable-ok: splatted into LiteLLM_TagTable, which declares this field as a dict
spend: float
budget_id: str | None
created_at: datetime | None
created_by: str | None
updated_at: datetime | None
class _TeamMembershipRowFields(TypedDict):
user_id: str
team_id: str
budget_id: str | None
spend: float | None
class _OrgMembershipRowFields(TypedDict):
user_id: str
organization_id: str
user_role: str | None
spend: float
budget_id: str | None
created_at: datetime
updated_at: datetime
user_email: str | None
class _UserRowFields(TypedDict):
user_id: str
user_alias: str | None
team_id: str | None
sso_user_id: str | None
organization_id: str | None
object_permission_id: str | None
teams: list[str] # mutable-ok: splatted into LiteLLM_UserTable, which declares this field as a list
user_role: str | None
max_budget: float | None
spend: float
user_email: str | None
models: list[str] # mutable-ok: splatted into LiteLLM_UserTable, which declares this field as a list
metadata: Optional[
dict[str, object]
] # mutable-ok: splatted into LiteLLM_UserTable, which declares this field as a dict
max_parallel_requests: int | None
tpm_limit: int | None
rpm_limit: int | None
budget_duration: str | None
budget_reset_at: datetime | None
allowed_cache_controls: list[
str
] # mutable-ok: splatted into LiteLLM_UserTable, which declares this field as a list
policies: list[str] # mutable-ok: splatted into LiteLLM_UserTable, which declares this field as a list
model_spend: Optional[
dict[str, object]
] # mutable-ok: splatted into LiteLLM_UserTable, which declares this field as a dict
model_max_budget: Optional[
dict[str, object]
] # mutable-ok: splatted into LiteLLM_UserTable, which declares this field as a dict
created_at: datetime | None
updated_at: datetime | None
class _TeamRowFields(TypedDict):
team_alias: str | None
team_id: str
organization_id: str | None
admins: list[str] # mutable-ok: splatted into LiteLLM_TeamTable, which declares this field as a list
members: list[str] # mutable-ok: splatted into LiteLLM_TeamTable, which declares this field as a list
team_member_permissions: Optional[
list[str]
] # mutable-ok: splatted into LiteLLM_TeamTable, which declares this field as a list
metadata: Optional[
dict[str, object]
] # mutable-ok: splatted into LiteLLM_TeamTable, which declares this field as a dict
tpm_limit: int | None
rpm_limit: int | None
max_budget: float | None
soft_budget: float | None
budget_duration: str | None
models: list[str] # mutable-ok: splatted into LiteLLM_TeamTable, which declares this field as a list
blocked: bool
router_settings: Optional[
dict[str, object]
] # mutable-ok: splatted into LiteLLM_TeamTable, which declares this field as a dict
access_group_ids: Optional[
list[str]
] # mutable-ok: splatted into LiteLLM_TeamTable, which declares this field as a list
default_team_member_models: Optional[
list[str]
] # mutable-ok: splatted into LiteLLM_TeamTable, which declares this field as a list
spend: float | None
max_parallel_requests: int | None
budget_reset_at: datetime | None
model_id: int | None
model_spend: Optional[
dict[str, object]
] # mutable-ok: splatted into LiteLLM_TeamTable, which declares this field as a dict
model_max_budget: Optional[
dict[str, object]
] # mutable-ok: splatted into LiteLLM_TeamTable, which declares this field as a dict
policies: list[str] | None # mutable-ok: splatted into LiteLLM_TeamTable, which declares this field as a list
allow_team_guardrail_config: bool | None
object_permission_id: str | None
updated_at: datetime | None
created_at: datetime | None
class _AccessGroupRowFields(TypedDict):
access_group_id: str
access_group_name: str
description: str | None
access_model_names: list[
str
] # mutable-ok: splatted into LiteLLM_AccessGroupTable, which declares this field as a list
access_mcp_server_ids: list[
str
] # mutable-ok: splatted into LiteLLM_AccessGroupTable, which declares this field as a list
access_agent_ids: list[
str
] # mutable-ok: splatted into LiteLLM_AccessGroupTable, which declares this field as a list
assigned_team_ids: list[
str
] # mutable-ok: splatted into LiteLLM_AccessGroupTable, which declares this field as a list
assigned_key_ids: list[
str
] # mutable-ok: splatted into LiteLLM_AccessGroupTable, which declares this field as a list
created_at: datetime | None
created_by: str | None
updated_at: datetime | None
updated_by: str | None
class _OrgRowFields(TypedDict):
organization_id: str | None
organization_alias: str | None
budget_id: str
spend: float
metadata: Optional[
dict[str, object]
] # mutable-ok: splatted into LiteLLM_OrganizationTable, which declares this field as a dict
models: list[str] # mutable-ok: splatted into LiteLLM_OrganizationTable, which declares this field as a list
model_spend: Optional[
dict[str, object]
] # mutable-ok: splatted into LiteLLM_OrganizationTable, which declares this field as a dict
created_by: str
updated_by: str
object_permission_id: str | None
class _ObjectPermissionRowFields(TypedDict):
object_permission_id: str
mcp_servers: Optional[
list[str]
] # mutable-ok: splatted into LiteLLM_ObjectPermissionTable, which declares this field as a list
mcp_access_groups: Optional[
list[str]
] # mutable-ok: splatted into LiteLLM_ObjectPermissionTable, which declares this field as a list
mcp_tool_permissions: Optional[
dict[str, list[str]]
] # mutable-ok: splatted into LiteLLM_ObjectPermissionTable, which declares this field as a dict
vector_stores: Optional[
list[str]
] # mutable-ok: splatted into LiteLLM_ObjectPermissionTable, which declares this field as a list
agents: Optional[
list[str]
] # mutable-ok: splatted into LiteLLM_ObjectPermissionTable, which declares this field as a list
agent_access_groups: Optional[
list[str]
] # mutable-ok: splatted into LiteLLM_ObjectPermissionTable, which declares this field as a list
models: Optional[
list[str]
] # mutable-ok: splatted into LiteLLM_ObjectPermissionTable, which declares this field as a list
mcp_toolsets: Optional[
list[str]
] # mutable-ok: splatted into LiteLLM_ObjectPermissionTable, which declares this field as a list
blocked_tools: Optional[
list[str]
] # mutable-ok: splatted into LiteLLM_ObjectPermissionTable, which declares this field as a list
search_tools: Optional[
list[str]
] # mutable-ok: splatted into LiteLLM_ObjectPermissionTable, which declares this field as a list
mcp_tool_search_enabled: bool | None
class _VectorStoreRowFields(TypedDict):
vector_store_id: str
custom_llm_provider: str
vector_store_name: str | None
vector_store_description: str | None
vector_store_metadata: Optional[
dict[str, object]
] # mutable-ok: splatted into LiteLLM_ManagedVectorStoresTable, which declares this field as a dict
created_at: datetime | None
updated_at: datetime | None
litellm_credential_name: str | None
litellm_params: Optional[
dict[str, object]
] # mutable-ok: splatted into LiteLLM_ManagedVectorStoresTable, which declares this field as a dict
team_id: str | None
user_id: str | None
class _ProjectRowFields(TypedDict):
project_id: str
project_alias: str | None
description: str | None
team_id: str | None
budget_id: str | None
metadata: Optional[
dict[str, object]
] # mutable-ok: splatted into LiteLLM_ProjectTableCachedObj, which declares this field as a dict
models: list[str] # mutable-ok: splatted into LiteLLM_ProjectTableCachedObj, which declares this field as a list
spend: float
model_spend: Optional[
dict[str, object]
] # mutable-ok: splatted into LiteLLM_ProjectTableCachedObj, which declares this field as a dict
model_rpm_limit: Optional[
dict[str, object]
] # mutable-ok: splatted into LiteLLM_ProjectTableCachedObj, which declares this field as a dict
model_tpm_limit: Optional[
dict[str, object]
] # mutable-ok: splatted into LiteLLM_ProjectTableCachedObj, which declares this field as a dict
blocked: bool
object_permission_id: str | None
created_by: str | None
updated_by: str | None
created_at: datetime | None
updated_at: datetime | None
class _UserAPIKeyAuthFields(TypedDict, total=False):
token: str | None
key_name: str | None
key_alias: str | None
spend: float
max_budget: float | None
user_id: str | None
team_id: str | None
models: list[str] # mutable-ok: splatted into UserAPIKeyAuth, which declares this field as a list
budget_duration: str | None
metadata: dict[str, object] # mutable-ok: splatted into UserAPIKeyAuth, which declares this field as a dict
class _OrgQueryKwargs(TypedDict, total=False):
where: Mapping[str, object]
include: Mapping[str, object]
class _BudgetRow(Protocol):
def dict(self) -> _BudgetRowFields: ...
class _EndUserRow(Protocol):
def dict(self) -> _EndUserRowFields: ...
class _TagRow(Protocol):
tag_name: str
def dict(self) -> _TagRowFields: ...
class _TeamMembershipRow(Protocol):
def dict(self) -> _TeamMembershipRowFields: ...
class _OrgMembershipRow(Protocol):
def model_dump(self) -> _OrgMembershipRowFields: ...
class _UserRow(Protocol):
user_id: str
@property
def organization_memberships(self) -> Sequence[_OrgMembershipRow | None] | None: ...
@organization_memberships.setter
def organization_memberships(self, value: Sequence[LiteLLM_OrganizationMembershipTable] | None) -> None: ...
def keys(self) -> Iterable[str]: ...
def __getitem__(self, key: str) -> object: ...
class _TeamRow(Protocol):
def dict(self) -> _TeamRowFields: ...
def model_dump(self) -> _TeamRowFields: ...
class _AccessGroupRow(Protocol):
def dict(self) -> _AccessGroupRowFields: ...
class _JWTKeyMappingRow(Protocol):
token: str
class _OrgRow(Protocol):
def model_dump(self) -> _OrgRowFields: ...
class _ObjectPermissionRow(Protocol):
def dict(self) -> _ObjectPermissionRowFields: ...
class _VectorStoreRow(Protocol):
def model_dump(self) -> _VectorStoreRowFields: ...
def dict(self) -> _VectorStoreRowFields: ...
def keys(self) -> Iterable[str]: ...
def __getitem__(self, key: str) -> object: ...
class _ProjectRow(Protocol):
def model_dump(self) -> _ProjectRowFields: ...
_PrismaRowT_co = TypeVar("_PrismaRowT_co", covariant=True)
class _PrismaTable(Protocol[_PrismaRowT_co]):
async def find_unique(
self,
where: Mapping[str, object] | None = None,
include: Mapping[str, object] | None = None,
) -> _PrismaRowT_co | None: ...
async def find_first(
self,
where: Mapping[str, object] | None = None,
include: Mapping[str, object] | None = None,
) -> _PrismaRowT_co | None: ...
async def find_many(
self,
where: Mapping[str, object] | None = None,
include: Mapping[str, object] | None = None,
take: int | None = None,
) -> Sequence[_PrismaRowT_co]: ...
async def create(
self,
data: Mapping[str, object],
include: Mapping[str, object] | None = None,
) -> _PrismaRowT_co: ...
async def update(
self,
where: Mapping[str, object],
data: Mapping[str, object],
) -> _PrismaRowT_co | None: ...
def _log_budget_lookup_failure(entity: str, error: Exception) -> None:
"""
Log a warning when budget lookup fails; cache will not be populated.
@ -377,7 +781,7 @@ _GUARDRAIL_MODIFICATION_KEYS: tuple = (
)
def _guardrail_modification_check(request_body: dict, team_object: Optional[LiteLLM_TeamTable]) -> None:
def _guardrail_modification_check(request_body: Mapping[str, object], team_object: LiteLLM_TeamTable | None) -> None:
"""
Reject user-supplied metadata flags that would modify guardrail behavior
unless the team has explicit permission. Checked keys include the plural
@ -392,7 +796,7 @@ def _guardrail_modification_check(request_body: dict, team_object: Optional[Lite
"""
from litellm.proxy.guardrails.guardrail_helpers import can_modify_guardrails
def _coerce_to_dict(container: Any) -> Optional[dict]:
def _coerce_to_dict(container: object) -> dict | None:
"""Accept dict or JSON-string (from multipart/form-data or extra_body).
Without this, an attacker can smuggle guardrail keys past the check by
@ -404,11 +808,11 @@ def _guardrail_modification_check(request_body: dict, team_object: Optional[Lite
if isinstance(container, dict):
return container
if isinstance(container, str):
parsed = safe_json_loads(container)
parsed = cast("object", safe_json_loads(container)) # cast-ok: safe_json_loads is untyped
return parsed if isinstance(parsed, dict) else None
return None
def _user_requested_modification(container: Any) -> bool:
def _user_requested_modification(container: object) -> bool:
coerced = _coerce_to_dict(container)
if coerced is None:
return False
@ -716,7 +1120,10 @@ async def common_checks(
_enforce_user_param_check(general_settings, request, request_body, route)
_global_proxy_budget_check(global_proxy_spend, skip_budget_checks, route)
_guardrail_modification_check(request_body, team_object)
_guardrail_modification_check(
cast("dict[str, object]", request_body), # cast-ok: request body dicts are string-keyed JSON payloads
team_object,
)
# 10 [OPTIONAL] Organization RBAC checks
organization_role_based_access_check(user_object=user_object, route=route, request_body=request_body)
@ -943,9 +1350,9 @@ async def get_default_end_user_budget(
# Fetch from database
try:
budget_record = await BudgetRepository(prisma_client).table.find_unique(
where={"budget_id": litellm.max_end_user_budget_id}
)
budget_record = await cast( # cast-ok: repository .table returns an untyped prisma client
"_PrismaTable[_BudgetRow]", BudgetRepository(prisma_client).table
).find_unique(where={"budget_id": litellm.max_end_user_budget_id})
if budget_record is None:
verbose_proxy_logger.warning(
@ -995,14 +1402,19 @@ async def get_team_member_default_budget(
cache_key = f"team_member_default_budget:{budget_id}"
cached_budget = await user_api_key_cache.async_get_cache(key=cache_key)
cached_budget = cast( # cast-ok: untyped cache returns the budget model or its serialized dict
"LiteLLM_BudgetTable | _BudgetRowFields | None",
await user_api_key_cache.async_get_cache(key=cache_key),
)
if isinstance(cached_budget, LiteLLM_BudgetTable):
return cached_budget
if isinstance(cached_budget, dict):
return LiteLLM_BudgetTable(**cached_budget)
try:
budget_record = await BudgetRepository(prisma_client).table.find_unique(where={"budget_id": budget_id})
budget_record = await cast( # cast-ok: repository .table returns an untyped prisma client
"_PrismaTable[_BudgetRow]", BudgetRepository(prisma_client).table
).find_unique(where={"budget_id": budget_id})
if budget_record is None:
verbose_proxy_logger.warning(f"Team-default member budget not found in database: {budget_id}")
@ -1159,7 +1571,9 @@ async def get_end_user_object(
# Fetch from database
try:
response = await EndUserRepository(prisma_client).table.find_unique(
response = await cast( # cast-ok: repository .table returns an untyped prisma client
"_PrismaTable[_EndUserRow]", EndUserRepository(prisma_client).table
).find_unique(
where={"user_id": end_user_id},
include={"litellm_budget_table": True, "object_permission": True},
)
@ -1231,7 +1645,9 @@ async def resolve_and_validate_end_user_id(
return raw_end_user_id
cache_key = f"end_user_validation:{raw_end_user_id}"
cached = await user_api_key_cache.async_get_cache(key=cache_key)
cached = cast( # cast-ok: untyped cache returns the validation marker string
"object", await user_api_key_cache.async_get_cache(key=cache_key)
)
if cached == "valid":
return raw_end_user_id
if cached == "invalid":
@ -1333,8 +1749,10 @@ async def get_tag_objects_batch(
if not tag_names:
return {}
tag_objects = {}
uncached_tags = []
tag_objects: dict[
str, LiteLLM_TagTable
] = {} # mutable-ok: filled per tag_name below and returned as the batch result
uncached_tags: list[str] = [] # mutable-ok: appended to in the cache-miss loop below
# Try to get all tags from cache first
for tag_name in tag_names:
@ -1351,7 +1769,9 @@ async def get_tag_objects_batch(
# Batch fetch uncached tags from DB in one query
if uncached_tags:
try:
db_tags = await TagRepository(prisma_client).table.find_many(
db_tags = await cast( # cast-ok: repository .table returns an untyped prisma client
"_PrismaTable[_TagRow]", TagRepository(prisma_client).table
).find_many(
where={"tag_name": {"in": uncached_tags}},
include={"litellm_budget_table": True},
)
@ -1445,7 +1865,9 @@ async def get_team_membership(
# else, check db
try:
response = await TeamMembershipRepository(prisma_client).table.find_unique(
response = await cast( # cast-ok: repository .table returns an untyped prisma client
"_PrismaTable[_TeamMembershipRow]", TeamMembershipRepository(prisma_client).table
).find_unique(
where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}},
include={"litellm_budget_table": True},
)
@ -1514,7 +1936,7 @@ def _should_check_db(key: str, last_db_access_time: LimitedSizeOrderedDict, db_c
return False
def _update_last_db_access_time(key: str, value: Optional[Any], last_db_access_time: LimitedSizeOrderedDict):
def _update_last_db_access_time(key: str, value: object | None, last_db_access_time: LimitedSizeOrderedDict):
last_db_access_time[key] = (value, time.time())
@ -1535,7 +1957,9 @@ def _get_role_based_permissions(
for role_based_permission in role_based_permissions:
if role_based_permission.role == rbac_role:
return getattr(role_based_permission, key)
return cast( # cast-ok: key is "models" or "routes", both Optional[List[str]] attributes
"list[str] | None", getattr(role_based_permission, key)
)
return None
@ -1576,7 +2000,7 @@ async def _get_fuzzy_user_object(
prisma_client: PrismaClient,
sso_user_id: Optional[str] = None,
user_email: Optional[str] = None,
) -> Optional[LiteLLM_UserTable]:
) -> Optional["_UserRow"]:
"""
Checks if sso user is in db.
@ -1590,7 +2014,9 @@ async def _get_fuzzy_user_object(
response = None
if sso_user_id is not None:
response = await UserRepository(prisma_client).table.find_unique(
response = await cast( # cast-ok: repository .table returns an untyped prisma client
"_PrismaTable[_UserRow]", UserRepository(prisma_client).table
).find_unique(
where={"sso_user_id": sso_user_id},
include={"organization_memberships": True},
)
@ -1598,14 +2024,18 @@ async def _get_fuzzy_user_object(
if response is None and user_email is not None:
# Use case-insensitive query to handle emails with different casing
# This matches the pattern used in _check_duplicate_user_email
response = await UserRepository(prisma_client).table.find_first(
response = await cast( # cast-ok: repository .table returns an untyped prisma client
"_PrismaTable[_UserRow]", UserRepository(prisma_client).table
).find_first(
where={"user_email": {"equals": user_email, "mode": "insensitive"}},
include={"organization_memberships": True},
)
if response is not None and sso_user_id is not None: # update sso_user_id
asyncio.create_task( # background task to update user with sso id
UserRepository(prisma_client).table.update(
cast( # cast-ok: repository .table returns an untyped prisma client
"_PrismaTable[_UserRow]", UserRepository(prisma_client).table
).update(
where={"user_id": response.user_id},
data={"sso_user_id": sso_user_id},
)
@ -1655,9 +2085,9 @@ async def get_user_object(
)
if should_check_db:
response = await UserRepository(prisma_client).table.find_unique(
where={"user_id": user_id}, include={"organization_memberships": True}
)
response = await cast( # cast-ok: repository .table returns an untyped prisma client
"_PrismaTable[_UserRow]", UserRepository(prisma_client).table
).find_unique(where={"user_id": user_id}, include={"organization_memberships": True})
if response is None:
response = await _get_fuzzy_user_object(
@ -1671,7 +2101,7 @@ async def get_user_object(
if response is None:
if user_id_upsert:
new_user_params: Dict[str, Any] = {
new_user_params: dict[str, object] = {
"user_id": user_id,
}
if user_email is not None:
@ -1683,10 +2113,14 @@ async def get_user_object(
and new_user_params.get("budget_reset_at") is None
):
new_user_params["budget_reset_at"] = get_budget_reset_time(
budget_duration=new_user_params["budget_duration"]
budget_duration=cast( # cast-ok: guarded non-None above; budget durations are strings
"str", new_user_params["budget_duration"]
)
)
response = await UserRepository(prisma_client).table.create(
response = await cast( # cast-ok: repository .table returns an untyped prisma client
"_PrismaTable[_UserRow]", UserRepository(prisma_client).table
).create(
data=new_user_params,
include={"organization_memberships": True},
)
@ -1708,7 +2142,7 @@ async def get_user_object(
]
response.organization_memberships = _dumped_memberships
_response = LiteLLM_UserTable(**dict(response))
_response = LiteLLM_UserTable(**cast("_UserRowFields", dict(response))) # cast-ok: prisma row field mapping
response_dict = _response.model_dump()
# save the user object to cache
@ -1736,7 +2170,7 @@ async def get_user_object(
async def _cache_management_object(
key: str,
value: Union[BaseModel, Dict[str, Any]],
value: BaseModel | dict[str, object],
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: Optional[ProxyLogging],
*,
@ -1830,7 +2264,9 @@ async def _delete_cache_key_object(
@log_db_metrics
async def _get_team_db_check(team_id: str, prisma_client: PrismaClient, team_id_upsert: Optional[bool] = None):
response = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id})
response = await cast( # cast-ok: repository .table returns an untyped prisma client
"_PrismaTable[_TeamRow]", TeamRepository(prisma_client).table
).find_unique(where={"team_id": team_id})
if response is None and team_id_upsert:
from litellm.proxy.management_endpoints.team_endpoints import new_team
@ -1845,12 +2281,17 @@ async def _get_team_db_check(team_id: str, prisma_client: PrismaClient, team_id_
http_request=mock_request,
user_api_key_dict=system_admin_user,
)
response = LiteLLM_TeamTable(**created_team_dict)
response = cast( # cast-ok: the validated team model exposes the same row surface
"_TeamRow",
LiteLLM_TeamTable(**cast("_TeamRowFields", created_team_dict)), # cast-ok: new_team returns an untyped dict
)
return response
async def _get_team_object_from_db(team_id: str, prisma_client: PrismaClient):
return await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id})
return await cast( # cast-ok: repository .table returns an untyped prisma client
"_PrismaTable[_TeamRow]", TeamRepository(prisma_client).table
).find_unique(where={"team_id": team_id})
async def _get_team_object_from_user_api_key_cache(
@ -2058,9 +2499,9 @@ async def get_access_object(
# Not in cache - fetch from DB
try:
response = await AccessGroupRepository(prisma_client).table.find_unique(
where={"access_group_id": access_group_id}
)
response = await cast( # cast-ok: repository .table returns an untyped prisma client
"_PrismaTable[_AccessGroupRow]", AccessGroupRepository(prisma_client).table
).find_unique(where={"access_group_id": access_group_id})
if response is None:
raise HTTPException(
@ -2134,7 +2575,9 @@ async def get_team_object_by_alias(
# Query database by team_alias
try:
teams = await TeamRepository(prisma_client).table.find_many(where={"team_alias": team_alias})
teams = await cast( # cast-ok: repository .table returns an untyped prisma client
"_PrismaTable[_TeamRow]", TeamRepository(prisma_client).table
).find_many(where={"team_alias": team_alias})
if not teams:
raise HTTPException(
@ -2236,7 +2679,9 @@ async def get_org_object_by_alias(
# Query database by organization_alias
try:
orgs = await OrganizationRepository(prisma_client).table.find_many(where={"organization_alias": org_alias})
orgs = await cast( # cast-ok: repository .table returns an untyped prisma client
"_PrismaTable[_OrgRow]", OrganizationRepository(prisma_client).table
).find_many(where={"organization_alias": org_alias})
if not orgs:
raise HTTPException(
@ -2396,7 +2841,9 @@ class ExperimentalUIJWTToken:
if decrypted_token is None:
return None
try:
return UserAPIKeyAuth(**json.loads(decrypted_token))
return UserAPIKeyAuth(
**cast("_UserAPIKeyAuthFields", json.loads(decrypted_token)) # cast-ok: decrypted UI session payload
)
except Exception as e:
raise Exception(f"Invalid hash key. Hash key={hashed_token}. Decrypted token={decrypted_token}. Error: {e}")
@ -2453,7 +2900,9 @@ async def get_jwt_key_mapping_object(
Returns the hashed token (str) if a matching active mapping is found, else None.
"""
mapping = await JWTKeyMappingRepository(prisma_client).table.find_first(
mapping = await cast( # cast-ok: repository .table returns an untyped prisma client
"_PrismaTable[_JWTKeyMappingRow]", JWTKeyMappingRepository(prisma_client).table
).find_first(
where={
"jwt_claim_name": jwt_claim_name,
"jwt_claim_value": jwt_claim_value,
@ -2515,7 +2964,9 @@ async def get_key_object(
code=status.HTTP_401_UNAUTHORIZED,
)
_response = UserAPIKeyAuth(**_valid_token.model_dump(exclude_none=True))
_response = UserAPIKeyAuth(
**cast("_UserAPIKeyAuthFields", _valid_token.model_dump(exclude_none=True)) # cast-ok: combined view dump
)
# Load object_permission if object_permission_id exists but object_permission is not loaded
if _response.object_permission_id and not _response.object_permission:
@ -2581,9 +3032,9 @@ async def get_object_permission(
# else, check db
try:
response = await ObjectPermissionRepository(prisma_client).table.find_unique(
where={"object_permission_id": object_permission_id}
)
response = await cast( # cast-ok: repository .table returns an untyped prisma client
"_PrismaTable[_ObjectPermissionRow]", ObjectPermissionRepository(prisma_client).table
).find_unique(where={"object_permission_id": object_permission_id})
if response is None:
return None
@ -2637,7 +3088,9 @@ async def get_managed_vector_store_rows_by_uuids(
if not cache_misses:
return result
rows = await ManagedVectorStoresRepository(prisma_client).table.find_many(
rows = await cast( # cast-ok: repository .table returns an untyped prisma client
"_PrismaTable[_VectorStoreRow]", ManagedVectorStoresRepository(prisma_client).table
).find_many(
where={"vector_store_id": {"in": cache_misses}},
take=len(cache_misses),
)
@ -2648,7 +3101,9 @@ async def get_managed_vector_store_rows_by_uuids(
row_dict = dict(row) if hasattr(row, "__dict__") else {}
if not row_dict:
continue
cached_obj = LiteLLM_ManagedVectorStoresTable(**row_dict)
cached_obj = LiteLLM_ManagedVectorStoresTable(
**cast("_VectorStoreRowFields", row_dict) # cast-ok: prisma row field mapping
)
key = "managed_vector_store_id:{}".format(cached_obj.vector_store_id)
await user_api_key_cache.async_set_cache(
key=key,
@ -2711,11 +3166,13 @@ async def get_org_object(
return deserialized_org
# else, check db
try:
query_kwargs: Dict[str, Any] = {"where": {"organization_id": org_id}}
query_kwargs: _OrgQueryKwargs = {"where": {"organization_id": org_id}}
if include_budget_table:
query_kwargs["include"] = {"litellm_budget_table": True}
response = await OrganizationRepository(prisma_client).table.find_unique(**query_kwargs)
response = await cast( # cast-ok: repository .table returns an untyped prisma client
"_PrismaTable[_OrgRow]", OrganizationRepository(prisma_client).table
).find_unique(**query_kwargs)
except Exception:
# An operational failure (DB down, timeout, cache fault) is NOT the same fact as a confirmed
# missing row, and relabelling it as "doesn't exist" made every caller unable to tell them
@ -3670,17 +4127,18 @@ async def _virtual_key_soft_budget_check(
)
def _parse_email_list(raw: Any) -> List[str]:
def _parse_email_list(raw: object) -> list[str]:
"""Parse emails from a list or comma-separated string."""
if isinstance(raw, list):
return [e.strip() for e in raw if isinstance(e, str) and e.strip()]
entries = cast("list[object]", raw) # cast-ok: narrowed to list; elements validated below
return [e.strip() for e in entries if isinstance(e, str) and e.strip()]
elif isinstance(raw, str):
return [e.strip() for e in raw.split(",") if e.strip()]
return []
def _normalize_alert_emails(
cfg: Optional[Dict[str, Any]],
cfg: Mapping[str, object] | None,
) -> Dict[str, List[str]]:
"""Coerce user-supplied threshold→recipients mapping to Dict[str, List[str]].
@ -3693,8 +4151,8 @@ def _normalize_alert_emails(
def _merge_budget_alert_email_configs(
global_cfg: Optional[Dict[str, Any]],
per_key_cfg: Optional[Dict[str, Any]],
global_cfg: Mapping[str, object] | None,
per_key_cfg: Mapping[str, object] | None,
) -> Optional[Dict[str, List[str]]]:
"""
Per-threshold additive merge: each threshold's recipient list is the union
@ -4197,7 +4655,9 @@ async def get_project_object(
return deserialized_project
# Fetch from DB
project_row = await ProjectRepository(prisma_client).table.find_unique(
project_row = await cast( # cast-ok: repository .table returns an untyped prisma client
"_PrismaTable[_ProjectRow]", ProjectRepository(prisma_client).table
).find_unique(
where={"project_id": project_id},
include={"litellm_budget_table": True},
)
@ -4498,7 +4958,9 @@ async def vector_store_access_check(
#########################################################
# Check if the key can access the vector store
if valid_token is not None and valid_token.object_permission_id is not None:
key_object_permission = await ObjectPermissionRepository(prisma_client).table.find_unique(
key_object_permission = await cast( # cast-ok: repository .table returns an untyped prisma client
"_PrismaTable[LiteLLM_ObjectPermissionTable]", ObjectPermissionRepository(prisma_client).table
).find_unique(
where={"object_permission_id": valid_token.object_permission_id},
)
if key_object_permission is not None:
@ -4510,7 +4972,9 @@ async def vector_store_access_check(
# Check if the team can access the vector store
if team_object is not None and team_object.object_permission_id is not None:
team_object_permission = await ObjectPermissionRepository(prisma_client).table.find_unique(
team_object_permission = await cast( # cast-ok: repository .table returns an untyped prisma client
"_PrismaTable[LiteLLM_ObjectPermissionTable]", ObjectPermissionRepository(prisma_client).table
).find_unique(
where={"object_permission_id": team_object.object_permission_id},
)
if team_object_permission is not None:

View file

@ -1,9 +1,22 @@
import asyncio
from datetime import datetime
from types import SimpleNamespace
from typing import Any, Awaitable, Callable, Dict, List, Optional, Set, Tuple, Union
from typing import (
Awaitable,
Callable,
Dict,
List,
Mapping,
Optional,
Protocol,
Set,
Tuple,
Union,
cast,
)
from fastapi import HTTPException, status
from typing_extensions import TypedDict
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import CommonProxyErrors
@ -16,6 +29,7 @@ from litellm.types.proxy.management_endpoints.common_daily_activity import (
BreakdownMetrics,
DailySpendData,
DailySpendMetadata,
GroupedData,
KeyMetadata,
KeyMetricWithMetadata,
MetricWithMetadata,
@ -33,8 +47,152 @@ _PRISMA_TO_PG_TABLE: Dict[str, str] = {
"litellm_dailytagspend": "LiteLLM_DailyTagSpend",
}
_WhereFilter = dict[str, str | list[str] | dict[str, list[str]]]
_WhereConditions = dict[str, str | _WhereFilter]
def update_metrics(existing_metrics: SpendMetrics, record: Any) -> SpendMetrics:
class _SpendMetricsRecord(Protocol):
@property
def spend(self) -> float | None: ...
@property
def prompt_tokens(self) -> int | None: ...
@property
def completion_tokens(self) -> int | None: ...
@property
def cache_read_input_tokens(self) -> int | None: ...
@property
def cache_creation_input_tokens(self) -> int | None: ...
@property
def compression_saved_tokens(self) -> int | None: ...
@property
def compression_savings_spend(self) -> float | None: ...
@property
def prompt_caching_savings_spend(self) -> float | None: ...
@property
def api_requests(self) -> int | None: ...
@property
def successful_requests(self) -> int | None: ...
@property
def failed_requests(self) -> int | None: ...
class _DailySpendRecord(_SpendMetricsRecord, Protocol):
@property
def spend(self) -> float: ...
@property
def date(self) -> str: ...
@property
def api_key(self) -> str: ...
@property
def model(self) -> str | None: ...
@property
def model_group(self) -> str | None: ...
@property
def custom_llm_provider(self) -> str | None: ...
@property
def mcp_namespaced_tool_name(self) -> str | None: ...
@property
def endpoint(self) -> str | None: ...
class _GroupingSetsRow(_SpendMetricsRecord, Protocol):
@property
def group_level(self) -> int: ...
@property
def date(self) -> str: ...
@property
def api_key(self) -> str | None: ...
@property
def model(self) -> str | None: ...
@property
def model_group(self) -> str | None: ...
@property
def custom_llm_provider(self) -> str | None: ...
@property
def mcp_namespaced_tool_name(self) -> str | None: ...
@property
def endpoint(self) -> str | None: ...
class _TokenMetadataRecord(Protocol):
@property
def token(self) -> str: ...
@property
def key_alias(self) -> str | None: ...
@property
def team_id(self) -> str | None: ...
class _TokenTable(Protocol):
def find_many(
self,
where: Mapping[str, object],
order: Mapping[str, str] | None = None,
) -> Awaitable[list[_TokenMetadataRecord]]: ...
class _TokenTableRepository(Protocol):
@property
def table(self) -> _TokenTable: ...
class _RawQueryDb(Protocol):
def query_raw(self, query: str, *params: str) -> Awaitable[list[dict[str, object]] | None]: ...
def _token_table(repository: _TokenTableRepository) -> _TokenTable:
return repository.table
class _DailySpendTable(Protocol):
def count(self, where: Mapping[str, object]) -> Awaitable[int]: ...
def find_many(
self,
where: Mapping[str, object],
order: list[Mapping[str, str]],
skip: int,
take: int,
) -> Awaitable[list[_DailySpendRecord]]: ...
class _KeyMetadataInfo(TypedDict, total=False):
key_alias: str | None
team_id: str | None
class _AggregatedSpendData(TypedDict):
results: list[DailySpendData]
totals: SpendMetrics
def update_metrics(existing_metrics: SpendMetrics, record: _SpendMetricsRecord) -> SpendMetrics:
"""Update metrics with new record data.
Rollup rows can carry None for numeric fields when SUM() spans zero rows
@ -66,19 +224,19 @@ def _is_user_agent_tag(tag: Optional[str]) -> bool:
return normalized_tag.startswith("user-agent:") or normalized_tag.startswith("user agent:")
def compute_tag_metadata_totals(records: List[Any]) -> SpendMetrics:
def compute_tag_metadata_totals(records: list[_DailySpendRecord]) -> SpendMetrics:
"""
Deduplicate spend metrics for tags using request_id, ignoring User-Agent prefixed tags.
Each unique request_id contributes at most one record (the tag with max spend) to metadata.
"""
deduped_records: Dict[str, Any] = {}
deduped_records: dict[str, _DailySpendRecord] = {}
for record in records:
request_id = getattr(record, "request_id", None)
request_id: str | None = getattr(record, "request_id", None)
if not request_id:
continue
tag_value = getattr(record, "tag", None)
tag_value: str | None = getattr(record, "tag", None)
if _is_user_agent_tag(tag_value):
continue
@ -94,12 +252,12 @@ def compute_tag_metadata_totals(records: List[Any]) -> SpendMetrics:
def update_breakdown_metrics(
breakdown: BreakdownMetrics,
record: Any,
model_metadata: Dict[str, Dict[str, Any]],
provider_metadata: Dict[str, Dict[str, Any]],
api_key_metadata: Dict[str, Dict[str, Any]],
record: _DailySpendRecord,
model_metadata: dict[str, dict[str, object]],
provider_metadata: dict[str, dict[str, object]],
api_key_metadata: dict[str, _KeyMetadataInfo],
entity_id_field: Optional[str] = None,
entity_metadata_field: Optional[Dict[str, dict]] = None,
entity_metadata_field: dict[str, dict[str, object]] | None = None,
) -> BreakdownMetrics:
"""Updates breakdown metrics for a single record using the existing update_metrics function"""
@ -241,7 +399,7 @@ def update_breakdown_metrics(
# Update entity-specific metrics if entity_id_field is provided
if entity_id_field:
entity_value = getattr(record, entity_id_field, None)
entity_value: str | None = getattr(record, entity_id_field, None)
entity_value = entity_value if entity_value else "Unassigned" # allow for null entity_id_field
if entity_value not in breakdown.entities:
breakdown.entities[entity_value] = MetricWithMetadata(
@ -270,22 +428,24 @@ def update_breakdown_metrics(
async def get_api_key_metadata(
prisma_client: PrismaClient,
api_keys: Set[str],
) -> Dict[str, Dict[str, Any]]:
) -> dict[str, _KeyMetadataInfo]:
"""Get api key metadata, falling back to deleted keys table for keys not found in active table.
This ensures that key_alias and team_id are preserved in historical activity logs
even after a key is deleted or regenerated.
"""
key_records = await VerificationTokenRepository(prisma_client).table.find_many(
key_records = await _token_table(VerificationTokenRepository(prisma_client)).find_many(
where={"token": {"in": list(api_keys)}}
)
result = {k.token: {"key_alias": k.key_alias, "team_id": k.team_id} for k in key_records}
result: dict[str, _KeyMetadataInfo] = {
k.token: {"key_alias": k.key_alias, "team_id": k.team_id} for k in key_records
}
# For any keys not found in the active table, check the deleted keys table
missing_keys = api_keys - set(result.keys())
if missing_keys:
try:
deleted_key_records = await DeletedVerificationTokenRepository(prisma_client).table.find_many(
deleted_key_records = await _token_table(DeletedVerificationTokenRepository(prisma_client)).find_many(
where={"token": {"in": list(missing_keys)}},
order={"deleted_at": "desc"},
)
@ -342,12 +502,12 @@ def _build_where_conditions(
api_key: Optional[Union[str, List[str]]],
exclude_entity_ids: Optional[List[str]] = None,
timezone_offset_minutes: Optional[int] = None,
) -> Dict[str, Any]:
) -> _WhereConditions:
"""Build prisma where clause for daily activity queries."""
# Adjust dates for timezone if provided
adjusted_start, adjusted_end = _adjust_dates_for_timezone(start_date, end_date, timezone_offset_minutes)
where_conditions: Dict[str, Any] = {
where_conditions: _WhereConditions = {
"date": {
"gte": adjusted_start,
"lte": adjusted_end,
@ -369,7 +529,7 @@ def _build_where_conditions(
where_conditions[entity_id_field] = {"equals": entity_id}
if exclude_entity_ids:
current = where_conditions.get(entity_id_field, {})
current: str | _WhereFilter = where_conditions.get(entity_id_field, {})
if isinstance(current, str):
current = {"equals": current}
current["not"] = {"in": exclude_entity_ids}
@ -389,7 +549,7 @@ def _build_aggregated_sql_query(
api_key: Optional[str],
exclude_entity_ids: Optional[List[str]] = None,
timezone_offset_minutes: Optional[int] = None,
) -> Tuple[str, List[Any]]:
) -> tuple[str, list[str]]:
"""Build a parameterized SQL GROUP BY query for aggregated daily activity.
Groups by (date, api_key, model, model_group, custom_llm_provider,
@ -407,7 +567,7 @@ def _build_aggregated_sql_query(
adjusted_start, adjusted_end = _adjust_dates_for_timezone(start_date, end_date, timezone_offset_minutes)
sql_conditions: List[str] = []
sql_params: List[Any] = []
sql_params: list[str] = []
p = 1 # parameter index (1-based for PostgreSQL $N placeholders)
# Date range (always present)
@ -506,17 +666,17 @@ def _build_aggregated_sql_query(
def _aggregate_spend_records_sync(
*,
records: List[Any],
api_key_metadata: Dict[str, Dict[str, Any]],
records: list[_DailySpendRecord],
api_key_metadata: dict[str, _KeyMetadataInfo],
entity_id_field: Optional[str],
entity_metadata_field: Optional[Dict[str, dict]],
) -> Dict[str, Any]:
model_metadata: Dict[str, Dict[str, Any]] = {}
provider_metadata: Dict[str, Dict[str, Any]] = {}
entity_metadata_field: dict[str, dict[str, object]] | None,
) -> _AggregatedSpendData:
model_metadata: dict[str, dict[str, object]] = {}
provider_metadata: dict[str, dict[str, object]] = {}
results: List[DailySpendData] = []
total_metrics = SpendMetrics()
grouped_data: Dict[str, Dict[str, Any]] = {}
grouped_data: dict[str, GroupedData] = {}
for record in records:
date_str = record.date
@ -557,10 +717,10 @@ def _aggregate_spend_records_sync(
async def _aggregate_spend_records(
*,
prisma_client: PrismaClient,
records: List[Any],
records: list[_DailySpendRecord],
entity_id_field: Optional[str],
entity_metadata_field: Optional[Dict[str, dict]],
) -> Dict[str, Any]:
entity_metadata_field: dict[str, dict[str, object]] | None,
) -> _AggregatedSpendData:
"""Aggregate rows into DailySpendData list and total metrics.
The per-row loop is offloaded to a worker thread via asyncio.to_thread so
@ -568,7 +728,7 @@ async def _aggregate_spend_records(
"""
api_keys: Set[str] = {record.api_key for record in records if record.api_key}
api_key_metadata: Dict[str, Dict[str, Any]] = {}
api_key_metadata: dict[str, _KeyMetadataInfo] = {}
if api_keys:
api_key_metadata = await get_api_key_metadata(prisma_client, api_keys)
@ -603,7 +763,7 @@ _GROUP_DATE_ENDPOINT = 62 # 0b0111110
_GROUP_DATE_ENDPOINT_API_KEY = 30 # 0b0011110
def _record_to_spend_metrics(record: Any) -> SpendMetrics:
def _record_to_spend_metrics(record: _SpendMetricsRecord) -> SpendMetrics:
"""Build a SpendMetrics directly from one already-aggregated rollup row.
SUM() over zero rows is SQL NULL, so rollup rows (notably the grand-total
@ -627,16 +787,16 @@ def _record_to_spend_metrics(record: Any) -> SpendMetrics:
)
def _key_metadata(api_key_metadata: Dict[str, Dict[str, Any]], api_key: str) -> KeyMetadata:
def _key_metadata(api_key_metadata: dict[str, _KeyMetadataInfo], api_key: str) -> KeyMetadata:
meta = api_key_metadata.get(api_key, {})
return KeyMetadata(key_alias=meta.get("key_alias"), team_id=meta.get("team_id"))
def _aggregate_grouping_sets_records_sync(
*,
records: List[Any],
api_key_metadata: Dict[str, Dict[str, Any]],
) -> Dict[str, Any]:
records: list[_GroupingSetsRow],
api_key_metadata: dict[str, _KeyMetadataInfo],
) -> _AggregatedSpendData:
"""Build the response from rollup rows produced by the GROUPING SETS query.
Each row carries a `group_level` bitmask (from Postgres GROUPING()) that
@ -645,10 +805,10 @@ def _aggregate_grouping_sets_records_sync(
summing in Python and no nested update_metrics calls.
"""
total_metrics = SpendMetrics()
grouped_data: Dict[str, Dict[str, Any]] = {}
grouped_data: dict[str, GroupedData] = {}
def ensure_date(date_str: str) -> Dict[str, Any]:
bucket = grouped_data.get(date_str)
def ensure_date(date_str: str) -> GroupedData:
bucket: GroupedData | None = grouped_data.get(date_str)
if bucket is None:
bucket = {"metrics": SpendMetrics(), "breakdown": BreakdownMetrics()}
grouped_data[date_str] = bucket
@ -753,12 +913,12 @@ def _aggregate_grouping_sets_records_sync(
async def _aggregate_grouping_sets_records(
*,
prisma_client: PrismaClient,
records: List[Any],
) -> Dict[str, Any]:
records: list[_GroupingSetsRow],
) -> _AggregatedSpendData:
"""Async wrapper: fetch api_key_metadata, then dispatch on a worker thread."""
api_keys: Set[str] = {r.api_key for r in records if r.api_key}
api_key_metadata: Dict[str, Dict[str, Any]] = {}
api_key_metadata: dict[str, _KeyMetadataInfo] = {}
if api_keys:
api_key_metadata = await get_api_key_metadata(prisma_client, api_keys)
@ -774,7 +934,7 @@ async def get_daily_activity(
table_name: str,
entity_id_field: str,
entity_id: Optional[Union[str, List[str]]],
entity_metadata_field: Optional[Dict[str, dict]],
entity_metadata_field: dict[str, dict[str, object]] | None,
start_date: Optional[str],
end_date: Optional[str],
model: Optional[str],
@ -782,9 +942,11 @@ async def get_daily_activity(
page: int,
page_size: int,
exclude_entity_ids: Optional[List[str]] = None,
metadata_metrics_func: Optional[Callable[[List[Any]], SpendMetrics]] = None,
metadata_metrics_func: Callable[[list[_DailySpendRecord]], SpendMetrics] | None = None,
timezone_offset_minutes: Optional[int] = None,
resolve_entity_metadata: Optional[Callable[[list[Any]], Awaitable[dict[str, dict]]]] = None,
resolve_entity_metadata: Optional[
Callable[[list[_DailySpendRecord]], Awaitable[dict[str, dict[str, object]]]]
] = None,
) -> SpendAnalyticsPaginatedResponse:
"""Common function to get daily activity for any entity type.
@ -818,8 +980,12 @@ async def get_daily_activity(
timezone_offset_minutes=timezone_offset_minutes,
)
table = cast( # cast-ok: prisma table accessor is dynamic
_DailySpendTable, getattr(prisma_client.db, table_name)
)
# Get total count for pagination
total_count = await getattr(prisma_client.db, table_name).count(where=where_conditions)
total_count = await table.count(where=where_conditions)
# Fetch paginated results.
# ``date`` alone is not a unique sort key -- a busy tenant has many
@ -831,7 +997,7 @@ async def get_daily_activity(
# total. Adding ``id`` (the row's UUID primary key, present on both
# LiteLLM_DailyUserSpend and LiteLLM_DailyTeamSpend) as a tiebreaker
# gives every page a stable cursor (#30164).
daily_spend_data = await getattr(prisma_client.db, table_name).find_many(
daily_spend_data = await table.find_many(
where=where_conditions,
order=[
{"date": "desc"},
@ -893,7 +1059,7 @@ async def get_daily_activity_aggregated(
table_name: str,
entity_id_field: str,
entity_id: Optional[Union[str, List[str]]],
entity_metadata_field: Optional[Dict[str, dict]],
entity_metadata_field: dict[str, dict[str, object]] | None,
start_date: Optional[str],
end_date: Optional[str],
model: Optional[str],
@ -935,11 +1101,17 @@ async def get_daily_activity_aggregated(
)
# Execute GROUPING SETS query — returns one row per rollup level.
rows = await prisma_client.db.query_raw(sql_query, *sql_params)
db = cast(_RawQueryDb, prisma_client.db) # cast-ok: query_raw is resolved dynamically on the prisma wrapper
rows = await db.query_raw(sql_query, *sql_params)
if rows is None:
rows = []
records = [SimpleNamespace(**row) for row in rows]
records = [
cast( # cast-ok: raw rollup rows expose the selected SQL columns as attributes
_GroupingSetsRow, SimpleNamespace(**row)
)
for row in rows
]
# The grouping-sets dispatcher places each row directly in its bucket
# using the row's GROUPING() bitmask. No Python-side summing needed.

View file

@ -16,7 +16,7 @@ import asyncio
import json
import traceback
from datetime import datetime, timezone
from typing import Any, Dict, List, Optional, Union, cast
from typing import Any, Dict, List, Mapping, Optional, Protocol, Sequence, TypeVar, Union, cast
import fastapi
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
@ -72,10 +72,100 @@ from litellm.types.proxy.management_endpoints.internal_user_endpoints import (
)
if TYPE_CHECKING:
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.proxy_server import PrismaClient
from litellm.proxy.utils import ProxyLogging
router = APIRouter()
_PrismaRowT_co = TypeVar("_PrismaRowT_co", covariant=True)
class _TeamTableRow(Protocol):
team_id: str
members_with_roles: str
def model_dump(self) -> Mapping[str, Any]: ...
class _OrgMembershipRow(Protocol):
user_id: str
organization_id: str | None
class _PrismaTableClient(Protocol[_PrismaRowT_co]):
async def find_first(self, *, where: Mapping[str, object]) -> _PrismaRowT_co | None: ...
async def find_unique(self, *, where: Mapping[str, object]) -> _PrismaRowT_co | None: ...
async def find_many(
self,
*,
where: Mapping[str, object] = ...,
skip: int = ...,
take: int = ...,
order: Mapping[str, str] = ...,
) -> Sequence[_PrismaRowT_co]: ...
async def count(self, *, where: Mapping[str, object] = ...) -> int: ...
async def create(self, *, data: Mapping[str, object]) -> _PrismaRowT_co: ...
async def update(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> _PrismaRowT_co | None: ...
async def update_many(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> int: ...
async def delete_many(self, *, where: Mapping[str, object]) -> int: ...
def _user_table(prisma_client: Optional["PrismaClient"]) -> _PrismaTableClient[LiteLLM_UserTable]:
return cast( # cast-ok: repository .table is typed Any at the prisma seam
_PrismaTableClient[LiteLLM_UserTable],
UserRepository(prisma_client).table,
)
def _team_table(prisma_client: Optional["PrismaClient"]) -> _PrismaTableClient[_TeamTableRow]:
return cast( # cast-ok: repository .table is typed Any at the prisma seam
_PrismaTableClient[_TeamTableRow],
TeamRepository(prisma_client).table,
)
def _verification_token_table(prisma_client: Optional["PrismaClient"]) -> _PrismaTableClient[object]:
return cast( # cast-ok: repository .table is typed Any at the prisma seam
_PrismaTableClient[object],
VerificationTokenRepository(prisma_client).table,
)
def _invitation_link_table(prisma_client: Optional["PrismaClient"]) -> _PrismaTableClient[object]:
return cast( # cast-ok: repository .table is typed Any at the prisma seam
_PrismaTableClient[object],
InvitationLinkRepository(prisma_client).table,
)
def _team_membership_table(prisma_client: Optional["PrismaClient"]) -> _PrismaTableClient[object]:
return cast( # cast-ok: repository .table is typed Any at the prisma seam
_PrismaTableClient[object],
TeamMembershipRepository(prisma_client).table,
)
def _org_membership_table(prisma_client: Optional["PrismaClient"]) -> _PrismaTableClient[_OrgMembershipRow]:
return cast( # cast-ok: repository .table is typed Any at the prisma seam
_PrismaTableClient[_OrgMembershipRow],
OrganizationMembershipRepository(prisma_client).table,
)
def _organization_table(prisma_client: Optional["PrismaClient"]) -> _PrismaTableClient[object]:
return cast( # cast-ok: repository .table is typed Any at the prisma seam
_PrismaTableClient[object],
OrganizationRepository(prisma_client).table,
)
def _hash_password_in_dict(data: dict) -> None:
"""Hash password field in-place if present."""
@ -128,7 +218,7 @@ def _update_internal_new_user_params(data_json: dict, data: NewUserRequest) -> d
async def _check_duplicate_user_field(
field_name: str,
field_value: Optional[str],
prisma_client: Any,
prisma_client: Optional["PrismaClient"],
*,
case_insensitive: bool = False,
label: Optional[str] = None,
@ -156,7 +246,7 @@ async def _check_duplicate_user_field(
if case_insensitive:
where_clause[field_name]["mode"] = "insensitive"
existing_user = await UserRepository(prisma_client).table.find_first(where=where_clause)
existing_user = await _user_table(prisma_client).find_first(where=where_clause)
if existing_user is not None:
existing_value = getattr(existing_user, field_name, value)
@ -167,7 +257,7 @@ async def _check_duplicate_user_field(
)
async def _check_duplicate_user_email(user_email: Optional[str], prisma_client: Any) -> None:
async def _check_duplicate_user_email(user_email: str | None, prisma_client: Optional["PrismaClient"]) -> None:
"""
Helper function to check if a user email already exists in the database.
"""
@ -180,7 +270,7 @@ async def _check_duplicate_user_email(user_email: Optional[str], prisma_client:
)
async def _check_duplicate_user_id(user_id: Optional[str], prisma_client: Any) -> None:
async def _check_duplicate_user_id(user_id: str | None, prisma_client: Optional["PrismaClient"]) -> None:
"""
Helper function to check if a user id already exists in the database.
"""
@ -641,15 +731,15 @@ def _enforce_user_info_access(user_id: Optional[str], user_api_key_dict: UserAPI
async def _get_user_info_teams(
prisma_client: Any,
prisma_client: "PrismaClient",
user_id: Optional[str],
user_info: Optional[Any],
user_info: object | None,
user_api_key_dict: UserAPIKeyAuth,
) -> tuple[list[Any], Optional[list[Any]]]:
) -> tuple[list[TeamListResponseObject], list[TeamListResponseObject] | None]:
"""Fetch and merge teams from membership + user.teams field."""
from litellm.proxy.management_endpoints.team_endpoints import list_team
team_list: list[Any] = []
team_list: list[TeamListResponseObject] = []
team_id_list: list[str] = []
teams_1 = await list_team(
@ -664,7 +754,7 @@ async def _get_user_info_teams(
team_list = teams_1
team_id_list = [team.team_id for team in teams_1]
teams_2: Optional[list[Any]] = None
teams_2: list[TeamListResponseObject] | None = None
target_team_ids = getattr(user_info, "teams", None)
if target_team_ids and isinstance(target_team_ids, list):
@ -711,10 +801,12 @@ def _redact_scim_enterprise_metadata(
def _build_user_info_response(
user_id: Optional[str],
user_info: Optional[Any],
user_info: Optional[
BaseModel | dict[str, Any]
], # mutable-ok: password and metadata are stripped from this dict in place below
keys: Optional[List[LiteLLM_VerificationToken]],
team_list: list[Any],
teams_1: Optional[list[Any]],
team_list: list[TeamListResponseObject],
teams_1: list[TeamListResponseObject] | None,
) -> UserInfoResponse:
"""Create UserInfoResponse while filtering sensitive fields."""
if user_info is None and keys is not None:
@ -839,8 +931,8 @@ async def _check_user_info_v2_access(
return None
# Helper: fetch the target user row (reused across branches)
async def _fetch_target_user():
return await UserRepository(prisma_client).table.find_unique(where={"user_id": target_user_id})
async def _fetch_target_user() -> LiteLLM_UserTable | None:
return await _user_table(prisma_client).find_unique(where={"user_id": target_user_id})
# Rule 1: Proxy admins — fetch and return the target row directly
if _user_has_admin_view(user_api_key_dict):
@ -853,9 +945,7 @@ async def _check_user_info_v2_access(
# Rule 3: Team admins can look up users in their teams
if user_api_key_dict.user_id is not None:
# Get caller's teams
caller_user = await UserRepository(prisma_client).table.find_unique(
where={"user_id": user_api_key_dict.user_id}
)
caller_user = await _user_table(prisma_client).find_unique(where={"user_id": user_api_key_dict.user_id})
if caller_user is not None and caller_user.teams:
# Fetch the target user ONCE, before the loop
target_user = await _fetch_target_user()
@ -863,9 +953,9 @@ async def _check_user_info_v2_access(
return None
# Get all teams the caller belongs to
teams = await TeamRepository(prisma_client).table.find_many(where={"team_id": {"in": caller_user.teams}})
teams = await _team_table(prisma_client).find_many(where={"team_id": {"in": caller_user.teams}})
for team in teams:
team_obj = LiteLLM_TeamTable(**team.model_dump())
team_obj = LiteLLM_TeamTable.model_validate(team.model_dump())
if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj):
# Check if target user is in this team
if team.team_id in (target_user.teams or []):
@ -989,7 +1079,10 @@ async def _get_user_info_for_proxy_admin(user_api_key_dict: UserAPIKeyAuth):
"Database not connected. Connect a database to your proxy - https://docs.litellm.ai/docs/simple_proxy#managing-auth---virtual-keys"
)
results = await prisma_client.db.query_raw(sql_query)
results = cast( # cast-ok: prisma query_raw is untyped at the wrapper seam
list[dict[str, Any]],
await prisma_client.db.query_raw(sql_query),
)
verbose_proxy_logger.debug("results_keys: %s", results)
@ -1121,7 +1214,7 @@ def _update_internal_user_params(
async def _schedule_user_update_audit_log(
response: Dict[str, Any],
response: Mapping[str, object],
existing_user_row: Optional[BaseModel],
litellm_changed_by: Optional[str],
user_api_key_dict: UserAPIKeyAuth,
@ -1132,9 +1225,9 @@ async def _schedule_user_update_audit_log(
if prisma_client is None:
return
try:
updated_user_row = await UserRepository(prisma_client).table.find_first(where={"user_id": response["user_id"]})
updated_user_row = await _user_table(prisma_client).find_first(where={"user_id": response["user_id"]})
if updated_user_row:
user_row_typed = LiteLLM_UserTable(**updated_user_row.model_dump(exclude_none=True))
user_row_typed = LiteLLM_UserTable.model_validate(updated_user_row.model_dump(exclude_none=True))
asyncio.create_task(
UserManagementEventHooks.create_internal_user_audit_log(
user_id=user_row_typed.user_id,
@ -1160,7 +1253,7 @@ def _check_user_update_authz(
raise HTTPException(status_code=403, detail="Only proxy admins can modify user roles.")
if existing_user_row is not None:
typed_row = LiteLLM_UserTable(**existing_user_row.model_dump(exclude_none=True))
typed_row = LiteLLM_UserTable.model_validate(existing_user_row.model_dump(exclude_none=True))
if not can_user_call_user_update(user_api_key_dict=user_api_key_dict, user_info=typed_row):
raise HTTPException(
status_code=403,
@ -1179,7 +1272,7 @@ def _check_user_update_authz(
async def _invalidate_user_spend_counter_if_changed(
non_default_values: dict[str, Any],
non_default_values: Mapping[str, object],
) -> None:
"""Invalidate the cross-pod spend counter after a direct ``spend`` change.
@ -1225,18 +1318,14 @@ async def _update_single_user_helper(
existing_user_row: Optional[BaseModel] = None
if user_request.user_id:
existing_user_row = await UserRepository(prisma_client).table.find_first(
where={"user_id": user_request.user_id}
)
existing_user_row = await _user_table(prisma_client).find_first(where={"user_id": user_request.user_id})
elif user_request.user_email:
existing_user_row = await UserRepository(prisma_client).table.find_first(
where={"user_email": user_request.user_email}
)
existing_user_row = await _user_table(prisma_client).find_first(where={"user_email": user_request.user_email})
_check_user_update_authz(user_request, user_api_key_dict, existing_user_row)
if existing_user_row is not None:
existing_user_row = LiteLLM_UserTable(**existing_user_row.model_dump(exclude_none=True))
existing_user_row = LiteLLM_UserTable.model_validate(existing_user_row.model_dump(exclude_none=True))
# Prevent budget self-escalation (GHSA-wvg4-6222-3q4r): non-admin callers
# must not be able to raise their own budget/spend fields.
@ -1585,7 +1674,7 @@ async def bulk_user_update(
detail="Only proxy admins can update all users at once.",
)
# Optimized path for updating all users directly in database
all_users_in_db = await UserRepository(prisma_client).table.find_many(order={"created_at": "desc"})
all_users_in_db = await _user_table(prisma_client).find_many(order={"created_at": "desc"})
if not all_users_in_db:
raise HTTPException(
@ -1617,7 +1706,7 @@ async def bulk_user_update(
try:
# Perform bulk database update
await UserRepository(prisma_client).table.update_many(
await _user_table(prisma_client).update_many(
where={},
data=non_default_values, # Update all users
)
@ -1700,9 +1789,9 @@ async def bulk_user_update(
async def get_user_key_counts(
prisma_client,
prisma_client: Optional["PrismaClient"],
user_ids: Optional[List[str]] = None,
):
) -> Mapping[str, int]:
"""
Helper function to get the count of keys for each user using Prisma's count method.
@ -1718,11 +1807,11 @@ async def get_user_key_counts(
if not user_ids or len(user_ids) == 0:
return {}
result = {}
result: dict[str, int] = {} # mutable-ok: one key assigned per user_id inside the loop below
# Get count for each user_id individually
for user_id in user_ids:
count = await VerificationTokenRepository(prisma_client).table.count(
count = await _verification_token_table(prisma_client).count(
where={
"user_id": user_id,
"OR": [
@ -1771,9 +1860,9 @@ def _validate_sort_params(sort_by: Optional[str], sort_order: str) -> Optional[D
async def _authorize_user_list_request(
user_api_key_dict: UserAPIKeyAuth,
organization_ids: Optional[str],
prisma_client: Any,
user_api_key_cache: Any,
proxy_logging_obj: Any,
prisma_client: "PrismaClient",
user_api_key_cache: "UserApiKeyCache",
proxy_logging_obj: "ProxyLogging",
) -> Optional[str]:
"""
Authorize the /user/list request and return the (possibly scoped) organization_ids string.
@ -1911,7 +2000,7 @@ async def get_users(
skip = (page - 1) * page_size
# Build where conditions based on provided parameters
where_conditions: Dict[str, Any] = {}
where_conditions: dict[str, object] = {}
if role:
where_conditions["user_role"] = role
@ -1959,7 +2048,7 @@ async def get_users(
_validate_sort_params(sort_by, sort_order) if sort_by is not None and isinstance(sort_by, str) else None
)
users = await UserRepository(prisma_client).table.find_many(
users: Sequence[LiteLLM_UserTable] | None = await _user_table(prisma_client).find_many(
where=where_conditions,
skip=skip,
take=page_size,
@ -1967,9 +2056,10 @@ async def get_users(
)
# Get total count of user rows
total_count = await UserRepository(prisma_client).table.count(where=where_conditions)
total_count = await _user_table(prisma_client).count(where=where_conditions)
# Get key count for each user
user_key_counts: Mapping[str, int]
if users is not None:
user_key_counts = await get_user_key_counts(prisma_client, [user.user_id for user in users])
else:
@ -1986,7 +2076,11 @@ async def get_users(
for user in users:
user_dump = user.model_dump()
user_dump["metadata"] = _redact_scim_enterprise_metadata(user_dump.get("metadata"))
user_list.append(LiteLLM_UserTableWithKeyCount(**user_dump, key_count=user_key_counts.get(user.user_id, 0)))
user_list.append(
LiteLLM_UserTableWithKeyCount.model_validate(
{**user_dump, "key_count": user_key_counts.get(user.user_id, 0)}
)
)
else:
user_list = []
@ -2056,10 +2150,10 @@ async def delete_user(
# loop an org-admin of org-A could delete users in org-B by supplying
# {"user_ids": [victim_in_org_B], "organization_id": "org-A"}.
caller_is_proxy_admin = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value
caller_admin_org_ids: set = set()
caller_admin_org_ids: set[str] = set()
if not caller_is_proxy_admin:
caller_memberships = (
await OrganizationMembershipRepository(prisma_client).table.find_many(
await _org_membership_table(prisma_client).find_many(
where={
"user_id": user_api_key_dict.user_id,
"user_role": LitellmUserRoles.ORG_ADMIN.value,
@ -2077,9 +2171,9 @@ async def delete_user(
# Batch-fetch target memberships once before the per-user loop. Avoids
# an N+1 DB call when delete_user is called with a large user_ids list.
target_org_ids_by_user: Dict[str, set] = {}
target_org_ids_by_user: dict[str, set[str]] = {}
if not caller_is_proxy_admin:
all_target_memberships = await OrganizationMembershipRepository(prisma_client).table.find_many(
all_target_memberships = await _org_membership_table(prisma_client).find_many(
where={"user_id": {"in": data.user_ids}}
)
for m in all_target_memberships:
@ -2089,7 +2183,7 @@ async def delete_user(
# check that all teams passed exist
for user_id in data.user_ids:
user_row = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_id})
user_row = await _user_table(prisma_client).find_unique(where={"user_id": user_id})
if user_row is None:
raise HTTPException(
@ -2141,11 +2235,11 @@ async def delete_user(
)
## CLEANUP MEMBERS_WITH_ROLES
fetch_all_teams = await TeamRepository(prisma_client).table.find_many(where={"team_id": {"in": user_row.teams}})
teams_to_update = []
fetch_all_teams = await _team_table(prisma_client).find_many(where={"team_id": {"in": user_row.teams}})
teams_to_update: list[_TeamTableRow] = [] # mutable-ok: appended to inside the membership-cleanup loop below
for team in fetch_all_teams:
is_member_in_team, new_team_members = _cleanup_members_with_roles(
existing_team_row=LiteLLM_TeamTable(**team.model_dump()),
existing_team_row=LiteLLM_TeamTable.model_validate(team.model_dump()),
data=TeamMemberDeleteRequest(
team_id=team.team_id,
user_id=user_row.user_id,
@ -2160,17 +2254,17 @@ async def delete_user(
## update teams
for team in teams_to_update:
await TeamRepository(prisma_client).table.update(
await _team_table(prisma_client).update(
where={"team_id": team.team_id},
data={"members_with_roles": team.members_with_roles},
)
# End of Audit logging
## DELETE ASSOCIATED KEYS
await VerificationTokenRepository(prisma_client).table.delete_many(where={"user_id": {"in": data.user_ids}})
await _verification_token_table(prisma_client).delete_many(where={"user_id": {"in": data.user_ids}})
## DELETE ASSOCIATED INVITATION LINKS
await InvitationLinkRepository(prisma_client).table.delete_many(
await _invitation_link_table(prisma_client).delete_many(
where={
"OR": [
{"user_id": {"in": data.user_ids}},
@ -2181,13 +2275,13 @@ async def delete_user(
)
## DELETE ASSOCIATED ORGANIZATION MEMBERSHIPS
await OrganizationMembershipRepository(prisma_client).table.delete_many(where={"user_id": {"in": data.user_ids}})
await _org_membership_table(prisma_client).delete_many(where={"user_id": {"in": data.user_ids}})
## DELETE ASSOCIATED TEAM MEMBERSHIPS
await TeamMembershipRepository(prisma_client).table.delete_many(where={"user_id": {"in": data.user_ids}})
await _team_membership_table(prisma_client).delete_many(where={"user_id": {"in": data.user_ids}})
## DELETE USERS
deleted_users = await UserRepository(prisma_client).table.delete_many(where={"user_id": {"in": data.user_ids}})
deleted_users = await _user_table(prisma_client).delete_many(where={"user_id": {"in": data.user_ids}})
return deleted_users
@ -2215,14 +2309,14 @@ async def add_internal_user_to_organization(
try:
# Check if organization_id exists
organization_row = await OrganizationRepository(prisma_client).table.find_unique(
organization_row = await _organization_table(prisma_client).find_unique(
where={"organization_id": organization_id}
)
if organization_row is None:
raise Exception(f"Organization not found, passed organization_id={organization_id}")
# Create a new organization membership entry
new_membership = await OrganizationMembershipRepository(prisma_client).table.create(
new_membership = await _org_membership_table(prisma_client).create(
data={
"user_id": user_id,
"organization_id": organization_id,
@ -2239,9 +2333,9 @@ async def add_internal_user_to_organization(
async def _resolve_org_filter_for_user_search(
user_api_key_dict: UserAPIKeyAuth,
team_id: Optional[str],
prisma_client: Any,
user_api_key_cache: Any,
proxy_logging_obj: Any,
prisma_client: "PrismaClient",
user_api_key_cache: "UserApiKeyCache",
proxy_logging_obj: "ProxyLogging",
) -> Optional[List[str]]:
"""
Return a list of org IDs to filter by, or ``None`` for no filter.
@ -2305,9 +2399,9 @@ async def _resolve_org_filter_for_user_search(
async def _resolve_team_org_filter(
user_api_key_dict: UserAPIKeyAuth,
team_id: str,
prisma_client: Any,
user_api_key_cache: Any,
proxy_logging_obj: Any,
prisma_client: "PrismaClient",
user_api_key_cache: "UserApiKeyCache",
proxy_logging_obj: "ProxyLogging",
) -> List[str]:
"""Look up the team and return its org as a filter list, or raise 403."""
from litellm.proxy.management_endpoints.common_utils import _is_user_team_admin
@ -2397,7 +2491,7 @@ async def ui_view_users(
skip = (page - 1) * page_size
# Build where conditions based on provided parameters
where_conditions: Dict[str, Any] = {}
where_conditions: dict[str, object] = {}
if user_id:
where_conditions["user_id"] = {
@ -2416,7 +2510,7 @@ async def ui_view_users(
where_conditions["organization_memberships"] = {"some": {"organization_id": {"in": org_filter_ids}}}
# Query users with pagination and filters
users: Optional[List[BaseModel]] = await UserRepository(prisma_client).table.find_many(
users: Sequence[LiteLLM_UserTable] | None = await _user_table(prisma_client).find_many(
where=where_conditions,
skip=skip,
take=page_size,
@ -2426,7 +2520,7 @@ async def ui_view_users(
if not users:
return []
return [LiteLLM_UserTableFiltered(**user.model_dump()) for user in users]
return [LiteLLM_UserTableFiltered.model_validate(user.model_dump()) for user in users]
except HTTPException:
raise
@ -2438,13 +2532,15 @@ async def ui_view_users(
# Using shared metric helper implementations from common_daily_activity
async def _resolve_user_email_metadata(prisma_client: "PrismaClient", records: list[Any]) -> dict[str, dict]:
async def _resolve_user_email_metadata(
prisma_client: Optional["PrismaClient"], records: Sequence[Any]
) -> dict[str, dict]:
"""Map each user_id on the page to its email/alias so the Usage dashboard can
label the 'Spend Per User' chart with the email instead of the raw UUID."""
user_ids = {record.user_id for record in records if getattr(record, "user_id", None)}
if not user_ids:
return {}
users = await UserRepository(prisma_client).table.find_many(where={"user_id": {"in": list(user_ids)}})
users = await _user_table(prisma_client).find_many(where={"user_id": {"in": list(user_ids)}})
return {user.user_id: {"user_email": user.user_email, "user_alias": user.user_alias} for user in users}

View file

@ -18,11 +18,12 @@ import os
import re
import secrets
import traceback
from collections.abc import Mapping
from collections.abc import Mapping, Sequence
from datetime import datetime, timedelta, timezone
from typing import Any, Callable, Dict, List, Literal, Optional, Tuple, cast
import fastapi
from typing_extensions import Never
import yaml
from fastapi import APIRouter, Depends, Header, HTTPException, Query, Request, status
@ -127,6 +128,7 @@ from litellm.types.proxy.management_endpoints.key_management_endpoints import (
from litellm.types.router import Deployment
from litellm.types.utils import (
BudgetConfig,
CredentialItem,
PersonalUIKeyGenerationConfig,
TeamUIKeyGenerationConfig,
)
@ -490,7 +492,7 @@ _NON_ADMIN_SAFE_ALLOWED_ROUTES_PRESETS = frozenset({"llm_api_routes", "info_rout
def _validate_caller_can_change_key_ownership(
data: Optional[BaseModel],
existing_key_row: Any,
existing_key_row: LiteLLM_VerificationToken,
user_api_key_dict: UserAPIKeyAuth,
) -> None:
"""
@ -889,7 +891,7 @@ async def _common_key_generation_helper(
)
new_budget = prisma_client.jsonify_object(budget_row.json(exclude_none=True))
_budget = await BudgetRepository(prisma_client).table.create(
_budget: LiteLLM_BudgetTable = await BudgetRepository(prisma_client).table.create(
data={
**new_budget, # type: ignore
"created_by": user_api_key_dict.user_id or litellm_proxy_admin_name,
@ -1095,7 +1097,7 @@ async def _common_key_generation_helper(
def _check_key_model_specific_limits(
keys: List[LiteLLM_VerificationToken],
keys: Sequence[LiteLLM_VerificationToken],
data: Union[GenerateKeyRequest, UpdateKeyRequest],
entity_rpm_limit: Optional[int],
entity_tpm_limit: Optional[int],
@ -1166,7 +1168,7 @@ def _check_key_model_specific_limits(
def _check_key_rpm_tpm_limits(
keys: List[LiteLLM_VerificationToken],
keys: Sequence[LiteLLM_VerificationToken],
data: Union[GenerateKeyRequest, UpdateKeyRequest],
entity_rpm_limit: Optional[int],
entity_tpm_limit: Optional[int],
@ -1204,7 +1206,7 @@ def _check_key_rpm_tpm_limits(
def check_team_key_model_specific_limits(
keys: List[LiteLLM_VerificationToken],
keys: Sequence[LiteLLM_VerificationToken],
team_table: LiteLLM_TeamTableCachedObj,
data: Union[GenerateKeyRequest, UpdateKeyRequest],
) -> None:
@ -1229,7 +1231,7 @@ def check_team_key_model_specific_limits(
def check_team_key_rpm_tpm_limits(
keys: List[LiteLLM_VerificationToken],
keys: Sequence[LiteLLM_VerificationToken],
team_table: LiteLLM_TeamTableCachedObj,
data: Union[GenerateKeyRequest, UpdateKeyRequest],
) -> None:
@ -1261,7 +1263,7 @@ async def _check_team_key_limits(
# calculate allocated tpm/rpm limit
# check if specified tpm/rpm limit is greater than allocated tpm/rpm limit
keys = await VerificationTokenRepository(prisma_client).table.find_many(
keys: Sequence[LiteLLM_VerificationToken] = await VerificationTokenRepository(prisma_client).table.find_many(
where={"team_id": team_table.team_id},
)
# Exclude the key being updated to avoid double-counting its limits.
@ -1331,7 +1333,7 @@ async def _check_project_key_limits(
def check_org_key_model_specific_limits(
keys: List[LiteLLM_VerificationToken],
keys: Sequence[LiteLLM_VerificationToken],
org_table: LiteLLM_OrganizationTable,
data: Union[GenerateKeyRequest, UpdateKeyRequest],
) -> None:
@ -1364,7 +1366,7 @@ def check_org_key_model_specific_limits(
def check_org_key_rpm_tpm_limits(
keys: List[LiteLLM_VerificationToken],
keys: Sequence[LiteLLM_VerificationToken],
org_table: LiteLLM_OrganizationTable,
data: Union[GenerateKeyRequest, UpdateKeyRequest],
) -> None:
@ -1405,7 +1407,7 @@ async def _validate_caller_can_assign_key_org(
detail="Cannot assign a key to an organization without a user_id on the caller's token",
)
user_row = await UserRepository(prisma_client).table.find_unique(
user_row: LiteLLM_UserTable | None = await UserRepository(prisma_client).table.find_unique(
where={"user_id": user_api_key_dict.user_id},
include={"organization_memberships": True},
)
@ -1443,7 +1445,7 @@ async def _check_org_key_limits(
# get all organization keys
# calculate allocated tpm/rpm limit
# check if specified tpm/rpm limit is greater than allocated tpm/rpm limit
keys = await VerificationTokenRepository(prisma_client).table.find_many(
keys: Sequence[LiteLLM_VerificationToken] = await VerificationTokenRepository(prisma_client).table.find_many(
where={"organization_id": org_table.organization_id},
)
# Exclude the key being updated to avoid double-counting its limits.
@ -1587,7 +1589,7 @@ async def generate_key_fn(
if user_custom_key_generate is not None:
if inspect.iscoroutinefunction(user_custom_key_generate):
result = await user_custom_key_generate(data) # type: ignore
result: Mapping[str, object] = await user_custom_key_generate(data)
else:
raise ValueError("user_custom_key_generate must be a coroutine")
decision = result.get("decision", True)
@ -1786,7 +1788,7 @@ async def generate_service_account_key_fn(
if user_custom_key_generate is not None:
if inspect.iscoroutinefunction(user_custom_key_generate):
result = await user_custom_key_generate(data) # type: ignore
result: Mapping[str, object] = await user_custom_key_generate(data)
else:
raise ValueError("user_custom_key_generate must be a coroutine")
decision = result.get("decision", True)
@ -1855,7 +1857,7 @@ def prepare_metadata_fields(data: BaseModel, non_default_values: dict, existing_
)
casted_metadata[reserved_field] = existing_value
data_json = data.model_dump(exclude_unset=True, exclude_none=True)
data_json: dict[str, object] = data.model_dump(exclude_unset=True, exclude_none=True)
try:
for k, v in data_json.items():
@ -2046,7 +2048,9 @@ async def _get_and_validate_existing_key(
hashed_token = _hash_token_if_needed(token=token)
existing_key_row = await VerificationTokenRepository(prisma_client).table.find_unique(where={"token": hashed_token})
existing_key_row: LiteLLM_VerificationToken | None = await VerificationTokenRepository(
prisma_client
).table.find_unique(where={"token": hashed_token})
if existing_key_row is None:
raise ProxyException(
@ -2065,11 +2069,11 @@ async def _process_single_key_update(
litellm_changed_by: Optional[str],
prisma_client: Optional[PrismaClient],
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: Any,
proxy_logging_obj: ProxyLogging,
llm_router: Optional[Router],
user_custom_key_update: Optional[Callable] = None,
existing_key_row: Optional[LiteLLM_VerificationToken] = None,
) -> Dict[str, Any]:
) -> dict[str, object]:
"""
Process a single key update with all validations and checks.
@ -2218,9 +2222,9 @@ async def _process_single_key_update(
async def _validate_mcp_servers_for_key_update(
data: "UpdateKeyRequest",
team_obj: Optional["LiteLLM_TeamTableCachedObj"],
existing_key_row: Any,
prisma_client: Any,
user_api_key_cache: Any,
existing_key_row: LiteLLM_VerificationToken,
prisma_client: PrismaClient | None,
user_api_key_cache: UserApiKeyCache,
is_proxy_admin: bool,
) -> Optional[ObjectPermissionDict]:
"""Validate MCP servers in object_permission against the effective team."""
@ -2255,12 +2259,12 @@ async def _validate_mcp_servers_for_key_update(
async def _validate_update_key_data(
data: UpdateKeyRequest,
existing_key_row: Any,
existing_key_row: LiteLLM_VerificationToken,
user_api_key_dict: UserAPIKeyAuth,
llm_router: Any,
llm_router: Router | None,
premium_user: bool,
prisma_client: Any,
user_api_key_cache: Any,
user_api_key_cache: UserApiKeyCache,
) -> None:
"""Validate permissions and constraints for key update."""
# Reject NaN/±inf spend before it can reach the DB / spend counter.
@ -2614,7 +2618,7 @@ async def update_key_fn(
# Custom key update hook
if user_custom_key_update is not None:
if inspect.iscoroutinefunction(user_custom_key_update):
result = await user_custom_key_update(data)
result: Mapping[str, object] = await user_custom_key_update(data)
else:
raise ValueError("user_custom_key_update must be a coroutine")
decision = result.get("decision", True)
@ -2892,7 +2896,7 @@ def _build_failed_team_key_update(
else:
error_message = str(exception)
key_info: Optional[Dict[str, Any]] = None
key_info: dict[str, object] | None = None
if existing_key_row is not None:
if hasattr(existing_key_row, "model_dump"):
key_info = existing_key_row.model_dump()
@ -2962,7 +2966,9 @@ async def bulk_update_team_keys(
# `blocked` is Boolean? with no default; `/key/generate` writes NULL. Prisma's `NOT`
# excludes NULLs, so explicitly OR `false` with `null` to include them.
now = datetime.now(timezone.utc)
existing_keys = await VerificationTokenRepository(prisma_client).table.find_many(
existing_keys: Sequence[LiteLLM_VerificationToken] = await VerificationTokenRepository(
prisma_client
).table.find_many(
where={
"team_id": data.team_id,
"AND": [
@ -2980,7 +2986,7 @@ async def bulk_update_team_keys(
"error": f"Team {data.team_id} has more than {MAX_BATCH_SIZE} keys. Use `key_ids` to update in batches of {MAX_BATCH_SIZE}."
},
)
requested_tokens = [row.token for row in existing_keys]
requested_tokens = [cast(str, row.token) for row in existing_keys] # cast-ok: DB rows always carry a token
else:
if data.key_ids is None or len(data.key_ids) == 0:
raise HTTPException(
@ -3050,7 +3056,7 @@ async def bulk_update_team_keys(
update_key_request = UpdateKeyRequest(
key=token,
team_id=data.team_id,
**update_field_dict,
**cast(dict[str, Never], update_field_dict), # cast-ok: pydantic validates the dumped fields
)
updated_key_info = await _process_single_key_update(
update_key_request=update_key_request,
@ -3367,7 +3373,9 @@ async def info_key_fn_v2(
# Resolve key_aliases to tokens so we never pass token=None (unbounded query)
tokens_to_query = list(data.keys) if data.keys else []
if data.key_aliases:
alias_rows = await VerificationTokenRepository(prisma_client).table.find_many(
alias_rows: Sequence[LiteLLM_VerificationToken] = await VerificationTokenRepository(
prisma_client
).table.find_many(
where={"key_alias": {"in": data.key_aliases}},
include={"litellm_budget_table": True},
)
@ -4039,7 +4047,7 @@ def _transform_verification_tokens_to_deleted_records(
keys: List[LiteLLM_VerificationToken],
user_api_key_dict: UserAPIKeyAuth,
litellm_changed_by: Optional[str] = None,
) -> List[Dict[str, Any]]:
) -> list[dict[str, object]]:
"""Transform verification tokens into deleted token records ready for persistence."""
if not keys:
return []
@ -4049,13 +4057,15 @@ def _transform_verification_tokens_to_deleted_records(
for key in keys:
key_payload = key.model_dump()
deleted_record = LiteLLM_DeletedVerificationToken(
**key_payload,
**cast(dict[str, Never], key_payload), # cast-ok: pydantic revalidates the dumped fields
deleted_at=deleted_at,
deleted_by=user_api_key_dict.user_id,
deleted_by_api_key=user_api_key_dict.api_key,
litellm_changed_by=litellm_changed_by,
)
record = deleted_record.model_dump()
record: dict[str, object] = (
deleted_record.model_dump()
) # mutable-ok: org_id/relation keys are popped and JSON fields rewritten in place below
# Map org_id to organization_id (model uses org_id, but schema expects organization_id)
org_id_value = record.pop("org_id", None)
@ -4090,7 +4100,7 @@ def _transform_verification_tokens_to_deleted_records(
async def _save_deleted_verification_token_records(
records: List[Dict[str, Any]],
records: list[dict[str, object]],
prisma_client: PrismaClient,
) -> None:
"""Save deleted verification token records to the database."""
@ -4124,9 +4134,9 @@ async def delete_key_aliases(
user_api_key_dict: UserAPIKeyAuth,
litellm_changed_by: Optional[str] = None,
) -> Tuple[Optional[Dict], List[LiteLLM_VerificationToken]]:
_keys_being_deleted = await VerificationTokenRepository(prisma_client).table.find_many(
where={"key_alias": {"in": key_aliases}}
)
_keys_being_deleted: Sequence[LiteLLM_VerificationToken] = await VerificationTokenRepository(
prisma_client
).table.find_many(where={"key_alias": {"in": key_aliases}})
tokens = [key.token for key in _keys_being_deleted]
return await delete_verification_tokens(
@ -4161,7 +4171,7 @@ async def _rotate_master_key(
from litellm.proxy.proxy_server import proxy_config
try:
models: Optional[List] = await ModelRepository(prisma_client).table.find_many()
models: list[LiteLLM_ProxyModelTable] | None = await ModelRepository(prisma_client).table.find_many()
except Exception:
models = None
# 2. process model table
@ -4171,14 +4181,16 @@ async def _rotate_master_key(
new_models = []
for model in decrypted_models:
new_model = await _add_model_to_db(
model_params=Deployment(**model),
model_params=Deployment(**cast(dict[str, Never], model)), # cast-ok: pydantic revalidates
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
new_encryption_key=new_master_key,
should_create_model_in_db=False,
)
if new_model:
_dumped = new_model.model_dump(exclude_none=True)
_dumped: dict[str, object] = new_model.model_dump(
exclude_none=True
) # mutable-ok: litellm_params/model_info are rewritten in place as prisma.Json
_dumped["litellm_params"] = prisma.Json(_dumped["litellm_params"]) # type: ignore[attr-defined]
_dumped["model_info"] = prisma.Json(_dumped["model_info"]) # type: ignore[attr-defined]
new_models.append(_dumped)
@ -4191,7 +4203,7 @@ async def _rotate_master_key(
)
# 3. process config table
try:
config = await ConfigRepository(prisma_client).table.find_many()
config: Sequence[LiteLLM_Config] | None = await ConfigRepository(prisma_client).table.find_many()
except Exception:
config = None
@ -4256,7 +4268,7 @@ async def _rotate_master_key(
# 5. process credentials table
try:
credentials = await CredentialsRepository(prisma_client).table.find_many()
credentials: Sequence[CredentialItem] | None = await CredentialsRepository(prisma_client).table.find_many()
except Exception:
credentials = None
if credentials:
@ -4270,7 +4282,9 @@ async def _rotate_master_key(
updated_patch=decrypted_cred,
new_encryption_key=new_master_key,
)
_cred_data = encrypted_cred.model_dump(exclude_none=True)
_cred_data: dict[str, object] = encrypted_cred.model_dump(
exclude_none=True
) # mutable-ok: credential_values/credential_info are rewritten in place as prisma.Json
if "credential_values" in _cred_data:
_cred_data["credential_values"] = prisma.Json( # type: ignore[attr-defined]
_cred_data["credential_values"]
@ -4520,7 +4534,7 @@ async def _execute_virtual_key_regeneration(
grace_period=data.grace_period if data else None,
)
updated_token = await VerificationTokenRepository(prisma_client).table.update(
updated_token: LiteLLM_VerificationToken | None = await VerificationTokenRepository(prisma_client).table.update(
where={"token": hashed_api_key},
data=update_data, # type: ignore
)
@ -4535,7 +4549,7 @@ async def _execute_virtual_key_regeneration(
proxy_logging_obj=proxy_logging_obj,
)
response = GenerateKeyResponse(**updated_token_dict)
response = GenerateKeyResponse(**cast(dict[str, Never], updated_token_dict)) # cast-ok: pydantic revalidates
asyncio.create_task(
KeyManagementEventHooks.async_key_rotated_hook(
data=data,
@ -4721,7 +4735,9 @@ async def regenerate_key_fn(
else:
hashed_api_key = hash_token(key)
_key_in_db = await VerificationTokenRepository(prisma_client).table.find_unique(
_key_in_db: LiteLLM_VerificationToken | None = await VerificationTokenRepository(
prisma_client
).table.find_unique(
where={"token": hashed_api_key},
)
if _key_in_db is None:
@ -4853,7 +4869,7 @@ async def _check_proxy_or_team_admin_for_key(
)
def _validate_reset_spend_value(reset_to: Any, key_in_db: LiteLLM_VerificationToken) -> float:
def _validate_reset_spend_value(reset_to: object, key_in_db: LiteLLM_VerificationToken) -> float:
if not isinstance(reset_to, (int, float)):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
@ -4925,7 +4941,9 @@ async def reset_key_spend_fn(
else:
hashed_api_key = hash_token(key)
_key_in_db = await VerificationTokenRepository(prisma_client).table.find_unique(
_key_in_db: LiteLLM_VerificationToken | None = await VerificationTokenRepository(
prisma_client
).table.find_unique(
where={"token": hashed_api_key},
include={"litellm_budget_table": True},
)
@ -4945,7 +4963,7 @@ async def reset_key_spend_fn(
user_api_key_cache=user_api_key_cache,
)
updated_key = await VerificationTokenRepository(prisma_client).table.update(
updated_key: LiteLLM_VerificationToken | None = await VerificationTokenRepository(prisma_client).table.update(
where={"token": hashed_api_key},
data={"spend": reset_to},
)
@ -5029,7 +5047,9 @@ async def validate_key_list_check(
code=status.HTTP_403_FORBIDDEN,
)
complete_user_info = LiteLLM_UserTable(**complete_user_info_db_obj.model_dump())
complete_user_info = LiteLLM_UserTable(
**cast(dict[str, Never], complete_user_info_db_obj.model_dump()) # cast-ok: pydantic revalidates
)
# internal user can only see their own keys
if user_id:
@ -5063,7 +5083,7 @@ async def validate_key_list_check(
if key_hash:
try:
key_info = await VerificationTokenRepository(prisma_client).table.find_unique(
key_info: LiteLLM_VerificationToken = await VerificationTokenRepository(prisma_client).table.find_unique(
where={"token": key_hash},
)
except Exception:
@ -5102,7 +5122,10 @@ async def _fetch_user_team_objects(
if teams is None:
return []
return [LiteLLM_TeamTable(**team.model_dump()) for team in teams]
return [
LiteLLM_TeamTable(**cast(dict[str, Never], team.model_dump())) # cast-ok: pydantic revalidates
for team in teams
]
def _get_admin_team_ids_from_objects(
@ -5370,8 +5393,8 @@ async def list_keys(
async def _apply_non_admin_alias_scope(
user_api_key_dict: UserAPIKeyAuth,
prisma_client: Any,
query_params: List[Any],
prisma_client: PrismaClient,
query_params: list[object],
where_parts: List[str],
) -> None:
"""Append SQL scope conditions so non-admin users only see aliases for
@ -5384,7 +5407,9 @@ async def _apply_non_admin_alias_scope(
# Look up the user's teams from the user table
user_teams: List[str] = []
if user_api_key_dict.user_id:
user_row = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_api_key_dict.user_id})
user_row: LiteLLM_UserTable | None = await UserRepository(prisma_client).table.find_unique(
where={"user_id": user_api_key_dict.user_id}
)
if user_row is not None:
user_teams = getattr(user_row, "teams", []) or []
@ -5442,7 +5467,7 @@ async def key_aliases(
# support column-level SELECT projection on find_many.
#
# $1 is always UI_SESSION_TOKEN_TEAM_ID (filters out UI session tokens).
query_params: List[Any] = [UI_SESSION_TOKEN_TEAM_ID]
query_params: list[object] = [UI_SESSION_TOKEN_TEAM_ID]
where_parts = [
"key_alias IS NOT NULL",
"key_alias != ''",
@ -5469,7 +5494,7 @@ async def key_aliases(
where_sql = " AND ".join(where_parts)
count_sql = f'SELECT COUNT(*) AS count FROM "LiteLLM_VerificationToken" WHERE {where_sql}'
count_rows = await prisma_client.db.query_raw(count_sql, *query_params)
count_rows: Sequence[Mapping[str, int]] = await prisma_client.db.query_raw(count_sql, *query_params)
total_count = int(count_rows[0]["count"]) if count_rows else 0
aliases_params = query_params + [size, (page - 1) * size]
@ -5482,7 +5507,7 @@ async def key_aliases(
f" ORDER BY key_alias ASC"
f" LIMIT ${limit_idx} OFFSET ${offset_idx}"
)
alias_rows = await prisma_client.db.query_raw(aliases_sql, *aliases_params)
alias_rows: Sequence[Mapping[str, str]] = await prisma_client.db.query_raw(aliases_sql, *aliases_params)
aliases: List[str] = [row["key_alias"] for row in alias_rows if row.get("key_alias")]
total_pages = -(-total_count // size) if total_count > 0 else 0
@ -5550,7 +5575,7 @@ def _validate_sort_params(sort_by: Optional[str], sort_order: str) -> Optional[D
return order_by
def _build_expires_where_clause(expires_filter: str, now: datetime) -> dict[str, Any]:
def _build_expires_where_clause(expires_filter: str, now: datetime) -> dict[str, object]:
if expires_filter == "expired":
return {"AND": [{"expires": {"not": None}}, {"expires": {"lt": now}}]}
return {"OR": [{"expires": None}, {"expires": {"gte": now}}]}
@ -5764,7 +5789,9 @@ async def _list_key_helper(
# Fetch keys with pagination
if use_deleted_table:
keys = await DeletedVerificationTokenRepository(prisma_client).table.find_many(
keys: Sequence[LiteLLM_VerificationToken] = await DeletedVerificationTokenRepository(
prisma_client
).table.find_many(
where=where, # type: ignore
skip=skip, # type: ignore
take=size, # type: ignore
@ -5797,7 +5824,7 @@ async def _list_key_helper(
# Get total count of keys
if use_deleted_table:
total_count = await DeletedVerificationTokenRepository(prisma_client).table.count(
total_count: int = await DeletedVerificationTokenRepository(prisma_client).table.count(
where=where # type: ignore
)
else:
@ -5811,13 +5838,15 @@ async def _list_key_helper(
total_pages = -(-total_count // size) # Ceiling division
# Fetch user information if expand includes "user"
user_map = {}
user_map: Mapping[str, LiteLLM_UserTable] = {}
if expand and "user" in expand:
user_ids = [key.user_id for key in keys if key.user_id]
created_by_ids = [key.created_by for key in keys if key.created_by]
all_ids = list(set(user_ids + created_by_ids)) # Remove duplicates
if all_ids:
users = await UserRepository(prisma_client).table.find_many(where={"user_id": {"in": all_ids}})
users: Sequence[LiteLLM_UserTable] = await UserRepository(prisma_client).table.find_many(
where={"user_id": {"in": all_ids}}
)
user_map = {user.user_id: user for user in users}
# Prepare response
@ -5851,9 +5880,12 @@ async def _list_key_helper(
if return_full_object is True or (expand and "user" in expand):
if use_deleted_table:
# Use deleted key type to preserve deleted_at, deleted_by, etc.
key_list.append(LiteLLM_DeletedVerificationToken(**key_dict))
key_list.append(
LiteLLM_DeletedVerificationToken(**cast(dict[str, Never], key_dict)) # cast-ok: revalidated
)
else:
key_list.append(UserAPIKeyAuth(**key_dict)) # Return full key object
# Return full key object
key_list.append(UserAPIKeyAuth(**cast(dict[str, Never], key_dict))) # cast-ok: revalidated
else:
_token = key_dict.get("token")
key_list.append(cast(str, _token)) # Return only the token
@ -5880,8 +5912,8 @@ def _get_condition_to_filter_out_ui_session_tokens() -> Dict[str, Any]:
async def _check_key_admin_access(
user_api_key_dict: UserAPIKeyAuth,
hashed_token: str,
prisma_client: Any,
hashed_token: str | None,
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
route: str,
) -> None:
@ -5900,7 +5932,9 @@ async def _check_key_admin_access(
return
# Look up the target key to find its team
target_key_row = await VerificationTokenRepository(prisma_client).table.find_unique(where={"token": hashed_token})
target_key_row: LiteLLM_VerificationToken | None = await VerificationTokenRepository(
prisma_client
).table.find_unique(where={"token": hashed_token})
if target_key_row is None:
raise HTTPException(
status_code=404,
@ -5996,7 +6030,9 @@ async def block_key(
)
# Check if the key exists before trying to block it
existing_record = await VerificationTokenRepository(prisma_client).table.find_unique(where={"token": hashed_token})
existing_record: LiteLLM_VerificationToken | None = await VerificationTokenRepository(
prisma_client
).table.find_unique(where={"token": hashed_token})
if existing_record is None:
raise ProxyException(
message="Key not found.",
@ -6026,7 +6062,7 @@ async def block_key(
)
)
record = await VerificationTokenRepository(prisma_client).table.update(
record: LiteLLM_VerificationToken | None = await VerificationTokenRepository(prisma_client).table.update(
where={"token": hashed_token},
data={"blocked": True}, # type: ignore
)
@ -6107,7 +6143,9 @@ async def unblock_key(
)
# Check if the key exists before trying to unblock it
existing_record = await VerificationTokenRepository(prisma_client).table.find_unique(where={"token": hashed_token})
existing_record: LiteLLM_VerificationToken | None = await VerificationTokenRepository(
prisma_client
).table.find_unique(where={"token": hashed_token})
if existing_record is None:
raise ProxyException(
message="Key not found.",
@ -6137,7 +6175,7 @@ async def unblock_key(
)
)
record = await VerificationTokenRepository(prisma_client).table.update(
record: LiteLLM_VerificationToken | None = await VerificationTokenRepository(prisma_client).table.update(
where={"token": hashed_token},
data={"blocked": False}, # type: ignore
)
@ -6392,7 +6430,7 @@ def _validate_key_alias_format(key_alias: Optional[str]) -> None:
async def _enforce_unique_key_alias(
key_alias: Optional[str],
prisma_client: Any,
prisma_client: PrismaClient | None,
existing_key_token: Optional[str] = None,
) -> None:
"""
@ -6408,7 +6446,7 @@ async def _enforce_unique_key_alias(
ProxyException: If key alias already exists on a different key
"""
if key_alias is not None and prisma_client is not None:
where_clause: dict[str, Any] = {"key_alias": key_alias}
where_clause: dict[str, object] = {"key_alias": key_alias}
if existing_key_token:
# Exclude the current key from the uniqueness check
where_clause["NOT"] = {"token": existing_key_token}

View file

@ -13,7 +13,16 @@ Endpoints for /organization operations
#### ORGANIZATION MANAGEMENT ####
from typing import Annotated, Any, Dict, List, Mapping, Optional, Tuple
from typing import (
Annotated,
List,
Mapping,
Optional,
Protocol,
Sequence,
Tuple,
TypeVar,
)
import fastapi
from fastapi import APIRouter, Depends, HTTPException, Request, status
@ -108,6 +117,91 @@ _STR_OBJECT_DICT_ADAPTER = TypeAdapter(dict[str, object])
_BUDGET_SETTABLE_FIELDS = frozenset(LiteLLM_BudgetTable.model_fields.keys()) - {"budget_id"}
_ORG_COLUMN_FIELDS = frozenset({"organization_alias", "models"})
_PrismaModelT_co = TypeVar("_PrismaModelT_co", bound=BaseModel, covariant=True)
class _TableClient(Protocol[_PrismaModelT_co]):
async def find_unique(
self,
where: Mapping[str, object],
include: Mapping[str, object] | None = None,
) -> _PrismaModelT_co | None: ...
async def find_many(
self,
where: Mapping[str, object] | None = None,
include: Mapping[str, object] | None = None,
) -> Sequence[_PrismaModelT_co]: ...
async def create(
self,
data: Mapping[str, object],
include: Mapping[str, object] | None = None,
) -> _PrismaModelT_co: ...
async def update(
self,
where: Mapping[str, object],
data: Mapping[str, object],
include: Mapping[str, object] | None = None,
) -> _PrismaModelT_co | None: ...
async def upsert(
self,
where: Mapping[str, object],
data: Mapping[str, object],
include: Mapping[str, object] | None = None,
) -> _PrismaModelT_co: ...
async def delete(
self,
where: Mapping[str, object],
include: Mapping[str, object] | None = None,
) -> _PrismaModelT_co | None: ...
async def delete_many(self, where: Mapping[str, object] | None = None) -> int: ...
class _OrganizationTx(Protocol):
@property
def litellm_objectpermissiontable(self) -> "_TableClient[LiteLLM_ObjectPermissionTable]": ...
@property
def litellm_budgettable(self) -> "_TableClient[LiteLLM_BudgetTableFull]": ...
@property
def litellm_organizationtable(self) -> "_TableClient[LiteLLM_OrganizationTableWithMembers]": ...
def _organization_table(prisma_client: PrismaClient) -> "_TableClient[LiteLLM_OrganizationTableWithMembers]":
return OrganizationRepository(prisma_client).table
def _organization_membership_table(
prisma_client: PrismaClient,
) -> "_TableClient[LiteLLM_OrganizationMembershipTable]":
return OrganizationMembershipRepository(prisma_client).table
def _user_table(prisma_client: PrismaClient) -> "_TableClient[LiteLLM_UserTable]":
return UserRepository(prisma_client).table
def _budget_table(prisma_client: PrismaClient) -> "_TableClient[LiteLLM_BudgetTableFull]":
return BudgetRepository(prisma_client).table
def _object_permission_table(prisma_client: PrismaClient) -> "_TableClient[LiteLLM_ObjectPermissionTable]":
return ObjectPermissionRepository(prisma_client).table
def _team_table(prisma_client: PrismaClient) -> "_TableClient[LiteLLM_TeamTable]":
return TeamRepository(prisma_client).table
def _verification_token_table(prisma_client: PrismaClient) -> "_TableClient[LiteLLM_VerificationToken]":
return VerificationTokenRepository(prisma_client).table
def build_budget_write_data(budget_updates: Mapping[str, object], updated_by: str) -> Mapping[str, object]:
"""
@ -263,10 +357,9 @@ async def new_organization(
if user_api_key_dict.user_id is not None:
try:
user_object = await UserRepository(prisma_client).table.find_unique(
where={"user_id": user_api_key_dict.user_id}
)
user_object_correct_type = LiteLLM_UserTable(**user_object.model_dump())
user_object = await _user_table(prisma_client).find_unique(where={"user_id": user_api_key_dict.user_id})
if user_object is not None:
user_object_correct_type = LiteLLM_UserTable.model_validate(user_object.model_dump())
except Exception:
pass
@ -279,19 +372,21 @@ async def new_organization(
budget_params = LiteLLM_BudgetTable.model_fields.keys()
# Only include Budget Params when creating an entry in litellm_budgettable
_json_data = data.json(exclude_none=True)
_json_data = _STR_OBJECT_DICT_ADAPTER.validate_python(data.json(exclude_none=True))
_budget_data = {k: v for k, v in _json_data.items() if k in budget_params}
budget_row = LiteLLM_BudgetTable(**_budget_data)
budget_row = LiteLLM_BudgetTable.model_validate(_budget_data)
new_budget = prisma_client.jsonify_object(budget_row.json(exclude_none=True))
new_budget = _STR_OBJECT_DICT_ADAPTER.validate_python(
prisma_client.jsonify_object(budget_row.json(exclude_none=True))
)
_budget = await BudgetRepository(prisma_client).table.create(
_budget = await _budget_table(prisma_client).create(
data={
**new_budget, # type: ignore
**new_budget,
"created_by": user_api_key_dict.user_id or litellm_proxy_admin_name,
"updated_by": user_api_key_dict.user_id or litellm_proxy_admin_name,
}
) # type: ignore
)
data.budget_id = _budget.budget_id
@ -318,11 +413,13 @@ async def new_organization(
for m in data.models:
await can_user_call_model(m, llm_router=llm_router, user_object=user_object_correct_type)
organization_row = LiteLLM_OrganizationTable(
**data.json(exclude_none=True),
object_permission_id=object_permission_id,
created_by=user_api_key_dict.user_id or litellm_proxy_admin_name,
updated_by=user_api_key_dict.user_id or litellm_proxy_admin_name,
organization_row = LiteLLM_OrganizationTable.model_validate(
{
**data.json(exclude_none=True),
"object_permission_id": object_permission_id,
"created_by": user_api_key_dict.user_id or litellm_proxy_admin_name,
"updated_by": user_api_key_dict.user_id or litellm_proxy_admin_name,
}
)
for field in LiteLLM_ManagementEndpoint_MetadataFields:
@ -333,11 +430,13 @@ async def new_organization(
value=getattr(data, field),
)
new_organization_row = prisma_client.jsonify_object(organization_row.json(exclude_none=True))
new_organization_row = _STR_OBJECT_DICT_ADAPTER.validate_python(
prisma_client.jsonify_object(organization_row.json(exclude_none=True))
)
verbose_proxy_logger.info(f"new_organization_row: {json.dumps(new_organization_row, indent=2)}")
response = await OrganizationRepository(prisma_client).table.create(
response = await _organization_table(prisma_client).create(
data={
**new_organization_row, # type: ignore
**new_organization_row,
},
include={"litellm_budget_table": True},
)
@ -382,7 +481,7 @@ async def get_organization_daily_activity(
# Restrict non-proxy-admins to only organizations where they are org_admin
if not _user_has_admin_view(user_api_key_dict):
memberships = await OrganizationMembershipRepository(prisma_client).table.find_many(
memberships = await _organization_membership_table(prisma_client).find_many(
where={"user_id": user_api_key_dict.user_id}
)
admin_org_ids = [m.organization_id for m in memberships if m.user_role == LitellmUserRoles.ORG_ADMIN.value]
@ -399,11 +498,19 @@ async def get_organization_daily_activity(
)
# Fetch organization aliases for metadata
where_condition = {}
where_condition: dict[
str, object
] = {} # mutable-ok: the organization_id filter is inserted below only when org ids were requested
if org_ids_list:
where_condition["organization_id"] = {"in": list(org_ids_list)}
org_aliases = await OrganizationRepository(prisma_client).table.find_many(where=where_condition)
org_alias_metadata = {o.organization_id: {"organization_alias": o.organization_alias} for o in org_aliases}
org_aliases = await _organization_table(prisma_client).find_many(where=where_condition)
org_alias_metadata: dict[
str, dict[str, object]
] = { # mutable-ok: get_daily_activity's entity_metadata_field parameter is typed Optional[Dict[str, Dict[str, object]]]
o.organization_id: {"organization_alias": o.organization_alias}
for o in org_aliases
if o.organization_id is not None
}
# Query daily activity for organizations
return await get_daily_activity(
@ -436,7 +543,7 @@ async def _set_object_permission(
return None
if data.object_permission is not None:
created_object_permission = await ObjectPermissionRepository(prisma_client).table.create(
created_object_permission = await _object_permission_table(prisma_client).create(
data=data.object_permission.model_dump(exclude_none=True),
)
del data.object_permission
@ -474,7 +581,7 @@ async def update_organization(
)
# Transform UI payload to expected format
raw_data = await request.json()
raw_data = _STR_OBJECT_DICT_ADAPTER.validate_python(await request.json())
raw_data_with_flat_budget_fields = handle_nested_budget_structure_in_organization_update_request(raw_data)
# Create validated data model
@ -510,22 +617,24 @@ async def update_organization(
prisma_client=prisma_client,
)
existing_organization_row = await OrganizationRepository(prisma_client).table.find_unique(
existing_organization_row = await _organization_table(prisma_client).find_unique(
where={"organization_id": data.organization_id},
)
if existing_organization_row is None:
raise ValueError(f"Organization not found for organization_id={data.organization_id}")
updated_organization_row_json = data.model_dump(exclude_none=True)
updated_organization_row_json = _STR_OBJECT_DICT_ADAPTER.validate_python(data.model_dump(exclude_none=True))
# Merge metadata from existing organization with updated metadata
if updated_organization_row_json.get("metadata") is not None:
existing_metadata = existing_organization_row.metadata or {}
updated_metadata = updated_organization_row_json.get("metadata", {})
existing_metadata = _STR_OBJECT_DICT_ADAPTER.validate_python(existing_organization_row.metadata or {})
updated_metadata = _STR_OBJECT_DICT_ADAPTER.validate_python(updated_organization_row_json.get("metadata", {}))
merged_metadata = _update_dictionary(existing_dict=existing_metadata.copy(), new_dict=updated_metadata)
updated_organization_row_json["metadata"] = merged_metadata
updated_organization_row = prisma_client.jsonify_object(updated_organization_row_json)
updated_organization_row = _STR_OBJECT_DICT_ADAPTER.validate_python(
prisma_client.jsonify_object(updated_organization_row_json)
)
if data.object_permission is not None:
updated_organization_row = await handle_update_object_permission(
data_json=updated_organization_row,
@ -534,12 +643,16 @@ async def update_organization(
# Handle budget updates if budget fields are provided
budget_fields = {
k: v for k, v in data.model_dump().items() if k in LiteLLM_BudgetTable.model_fields.keys() and v is not None
k: v
for k, v in _STR_OBJECT_DICT_ADAPTER.validate_python(data.model_dump()).items()
if k in LiteLLM_BudgetTable.model_fields.keys() and v is not None
}
if budget_fields and existing_organization_row.budget_id:
await update_budget(
budget_obj=BudgetNewRequest(budget_id=existing_organization_row.budget_id, **budget_fields),
budget_obj=BudgetNewRequest.model_validate(
{"budget_id": existing_organization_row.budget_id, **budget_fields}
),
user_api_key_dict=user_api_key_dict,
)
@ -547,7 +660,7 @@ async def update_organization(
for field in LiteLLM_BudgetTable.model_fields.keys():
updated_organization_row.pop(field, None)
response = await OrganizationRepository(prisma_client).table.update(
response = await _organization_table(prisma_client).update(
where={"organization_id": data.organization_id},
data=updated_organization_row,
include={"members": True, "teams": True, "litellm_budget_table": True},
@ -557,9 +670,9 @@ async def update_organization(
async def handle_update_object_permission(
data_json: dict,
data_json: dict[str, object],
existing_organization_row: LiteLLM_OrganizationTable,
) -> dict:
) -> dict[str, object]:
"""
Handle the update of object permission for an organization.
@ -665,7 +778,7 @@ async def update_organization_v2(
prisma_client=prisma_client,
)
existing_organization_row = await OrganizationRepository(prisma_client).table.find_unique(
existing_organization_row = await _organization_table(prisma_client).find_unique(
where={"organization_id": organization_id},
)
if existing_organization_row is None:
@ -698,14 +811,17 @@ async def update_organization_v2(
else ({"object_permission_id": None} if object_permission_cleared else {})
)
organization_write_data = prisma_client.jsonify_object(
{
**org_column_updates,
**object_permission_write,
"updated_by": user_api_key_dict.user_id,
}
organization_write_data = _STR_OBJECT_DICT_ADAPTER.validate_python(
prisma_client.jsonify_object(
{
**org_column_updates,
**object_permission_write,
"updated_by": user_api_key_dict.user_id,
}
)
)
tx: _OrganizationTx
async with prisma_client.db.tx() as tx:
if object_permission_upsert is not None:
await tx.litellm_objectpermissiontable.upsert(
@ -718,8 +834,10 @@ async def update_organization_v2(
if budget_updates:
await tx.litellm_budgettable.update(
where={"budget_id": existing_organization_row.budget_id},
data=prisma_client.jsonify_object(
dict(build_budget_write_data(budget_updates, user_api_key_dict.user_id))
data=_STR_OBJECT_DICT_ADAPTER.validate_python(
prisma_client.jsonify_object(
dict(build_budget_write_data(budget_updates, user_api_key_dict.user_id))
)
),
)
response = await tx.litellm_organizationtable.update(
@ -762,18 +880,18 @@ async def delete_organization(
detail={"error": "Only proxy admins can delete organizations"},
)
deleted_orgs = []
deleted_orgs: list[
LiteLLM_OrganizationTableWithMembers
] = [] # mutable-ok: each deleted organization is appended by the loop below
for organization_id in data.organization_ids:
# delete all teams in the organization
await TeamRepository(prisma_client).table.delete_many(where={"organization_id": organization_id})
await _team_table(prisma_client).delete_many(where={"organization_id": organization_id})
# delete all members in the organization
await OrganizationMembershipRepository(prisma_client).table.delete_many(
where={"organization_id": organization_id}
)
await _organization_membership_table(prisma_client).delete_many(where={"organization_id": organization_id})
# delete all keys in the organization
await VerificationTokenRepository(prisma_client).table.delete_many(where={"organization_id": organization_id})
await _verification_token_table(prisma_client).delete_many(where={"organization_id": organization_id})
# delete the organization
deleted_org = await OrganizationRepository(prisma_client).table.delete(
deleted_org = await _organization_table(prisma_client).delete(
where={"organization_id": organization_id},
include={"members": True, "teams": True, "litellm_budget_table": True},
)
@ -836,7 +954,9 @@ async def list_organization(
)
# Build where conditions based on provided filters
where_conditions: Dict[str, Any] = {}
where_conditions: dict[
str, object
] = {} # mutable-ok: org_id / org_alias filter keys are inserted below as the query params arrive
if org_id:
where_conditions["organization_id"] = org_id
@ -848,14 +968,15 @@ async def list_organization(
}
# if proxy admin or admin viewer - get all orgs (with optional filters)
response: Sequence[LiteLLM_OrganizationTableWithMembers]
if _user_has_admin_view(user_api_key_dict):
response = await OrganizationRepository(prisma_client).table.find_many(
response = await _organization_table(prisma_client).find_many(
where=where_conditions if where_conditions else None,
include={"litellm_budget_table": True, "members": True, "teams": True},
)
# if internal user - get orgs they are a member of (with optional filters)
else:
org_memberships = await OrganizationMembershipRepository(prisma_client).table.find_many(
org_memberships = await _organization_membership_table(prisma_client).find_many(
where={"user_id": user_api_key_dict.user_id}
)
membership_org_ids = [membership.organization_id for membership in org_memberships]
@ -869,7 +990,7 @@ async def list_organization(
response = []
else:
where_conditions["organization_id"] = org_id
response = await OrganizationRepository(prisma_client).table.find_many(
response = await _organization_table(prisma_client).find_many(
where=where_conditions,
include={
"litellm_budget_table": True,
@ -880,7 +1001,7 @@ async def list_organization(
else:
# Filter by membership and any additional filters
where_conditions["organization_id"] = {"in": membership_org_ids}
response = await OrganizationRepository(prisma_client).table.find_many(
response = await _organization_table(prisma_client).find_many(
where=where_conditions,
include={
"litellm_budget_table": True,
@ -920,9 +1041,7 @@ async def info_organization(
prisma_client=prisma_client,
)
response: Optional[LiteLLM_OrganizationTableWithMembers] = await OrganizationRepository(
prisma_client
).table.find_unique(
response: LiteLLM_OrganizationTableWithMembers | None = await _organization_table(prisma_client).find_unique(
where={"organization_id": organization_id},
include={
"litellm_budget_table": True,
@ -939,7 +1058,7 @@ async def info_organization(
if response is None:
raise HTTPException(status_code=404, detail={"error": "Organization not found"})
response_pydantic_obj = LiteLLM_OrganizationTableWithMembers(**response.model_dump())
response_pydantic_obj = LiteLLM_OrganizationTableWithMembers.model_validate(response.model_dump())
return response_pydantic_obj
@ -975,7 +1094,7 @@ async def deprecated_info_organization(
prisma_client=prisma_client,
)
response = await OrganizationRepository(prisma_client).table.find_many(
response = await _organization_table(prisma_client).find_many(
where={"organization_id": {"in": data.organizations}},
include={"litellm_budget_table": True},
)
@ -1052,7 +1171,7 @@ async def organization_member_add(
)
# Check if organization exists
existing_organization_row = await OrganizationRepository(prisma_client).table.find_unique(
existing_organization_row = await _organization_table(prisma_client).find_unique(
where={"organization_id": data.organization_id}
)
if existing_organization_row is None:
@ -1125,7 +1244,7 @@ async def find_member_if_email(user_email: str, prisma_client: PrismaClient) ->
"error": f"Unique user not found for user_email={user_email}. Potential duplicate OR non-existent user_email in LiteLLM_UserTable. Use 'user_id' instead."
},
)
existing_user_email_row_pydantic = LiteLLM_UserTable(**existing_user_email_row.model_dump())
existing_user_email_row_pydantic = LiteLLM_UserTable.model_validate(existing_user_email_row.model_dump())
return existing_user_email_row_pydantic
@ -1163,7 +1282,7 @@ async def organization_member_update(
)
# Check if organization exists
existing_organization_row = await OrganizationRepository(prisma_client).table.find_unique(
existing_organization_row = await _organization_table(prisma_client).find_unique(
where={"organization_id": data.organization_id}
)
if existing_organization_row is None:
@ -1180,7 +1299,7 @@ async def organization_member_update(
data.user_id = existing_user_email_row.user_id
try:
existing_organization_membership = await OrganizationMembershipRepository(prisma_client).table.find_unique(
existing_organization_membership = await _organization_membership_table(prisma_client).find_unique(
where={
"user_id_organization_id": {
"user_id": data.user_id,
@ -1205,7 +1324,7 @@ async def organization_member_update(
# org-scoped operations. An org-admin of any org could otherwise
# alter a PROXY_ADMIN user's per-org role, which has downstream
# effects on admin UI filtering and scope derivation.
target_user_row = await UserRepository(prisma_client).table.find_unique(where={"user_id": data.user_id})
target_user_row = await _user_table(prisma_client).find_unique(where={"user_id": data.user_id})
if target_user_row is not None and getattr(target_user_row, "user_role", None) in (
LitellmUserRoles.PROXY_ADMIN.value,
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value,
@ -1222,7 +1341,7 @@ async def organization_member_update(
# Update member role
if data.role is not None:
await OrganizationMembershipRepository(prisma_client).table.update(
await _organization_membership_table(prisma_client).update(
where={
"user_id_organization_id": {
"user_id": data.user_id,
@ -1245,7 +1364,7 @@ async def organization_member_update(
)
# update organization membership with new budget_id
await OrganizationMembershipRepository(prisma_client).table.update(
await _organization_membership_table(prisma_client).update(
where={
"user_id_organization_id": {
"user_id": data.user_id,
@ -1254,9 +1373,7 @@ async def organization_member_update(
},
data={"budget_id": budget_id},
)
final_organization_membership: Optional[BaseModel] = await OrganizationMembershipRepository(
prisma_client
).table.find_unique(
final_organization_membership = await _organization_membership_table(prisma_client).find_unique(
where={
"user_id_organization_id": {
"user_id": data.user_id,
@ -1272,8 +1389,8 @@ async def organization_member_update(
detail={"error": f"Member not found in organization={data.organization_id} for user_id={data.user_id}"},
)
final_organization_membership_pydantic = LiteLLM_OrganizationMembershipTable(
**final_organization_membership.model_dump(exclude_none=True)
final_organization_membership_pydantic = LiteLLM_OrganizationMembershipTable.model_validate(
final_organization_membership.model_dump(exclude_none=True)
)
return final_organization_membership_pydantic
except Exception as e:
@ -1315,7 +1432,7 @@ async def organization_member_delete(
existing_user_email_row = await find_member_if_email(data.user_email, prisma_client)
data.user_id = existing_user_email_row.user_id
member_to_delete = await OrganizationMembershipRepository(prisma_client).table.delete(
member_to_delete = await _organization_membership_table(prisma_client).delete(
where={
"user_id_organization_id": {
"user_id": data.user_id,
@ -1374,16 +1491,16 @@ async def add_member_to_organization(
_returned_user = await prisma_client.insert_data(data=new_user_defaults, table_name="user") # type: ignore
if _returned_user is not None:
user_object = LiteLLM_UserTable(**_returned_user.model_dump())
user_object = LiteLLM_UserTable.model_validate(_returned_user.model_dump())
elif existing_user_email_row is not None and len(existing_user_email_row) > 1:
raise HTTPException(
status_code=400,
detail={"error": "Multiple users with this email found in db. Please use 'user_id' instead."},
)
elif existing_user_email_row is not None:
user_object = LiteLLM_UserTable(**existing_user_email_row.model_dump())
user_object = LiteLLM_UserTable.model_validate(existing_user_email_row.model_dump())
elif existing_user_id_row is not None:
user_object = LiteLLM_UserTable(**existing_user_id_row.model_dump())
user_object = LiteLLM_UserTable.model_validate(existing_user_id_row.model_dump())
else:
raise HTTPException(
status_code=404,
@ -1396,14 +1513,16 @@ async def add_member_to_organization(
)
# Add user to organization
_organization_membership = await OrganizationMembershipRepository(prisma_client).table.create(
_organization_membership = await _organization_membership_table(prisma_client).create(
data={
"organization_id": organization_id,
"user_id": user_object.user_id,
"user_role": member.role,
}
)
organization_membership = LiteLLM_OrganizationMembershipTable(**_organization_membership.model_dump())
organization_membership = LiteLLM_OrganizationMembershipTable.model_validate(
_organization_membership.model_dump()
)
return user_object, organization_membership
except Exception as e:

File diff suppressed because it is too large Load diff

View file

@ -15,7 +15,7 @@ import threading
import time
import traceback
import warnings
from collections.abc import Mapping
from collections.abc import Mapping, Sequence
from datetime import datetime, timedelta, timezone
from typing import (
TYPE_CHECKING,
@ -635,7 +635,7 @@ from fastapi.responses import (
RedirectResponse,
StreamingResponse,
)
from fastapi.routing import APIRouter
from fastapi.routing import APIRoute, APIRouter
from fastapi.security import OAuth2PasswordBearer
from fastapi.security.api_key import APIKeyHeader
from fastapi.staticfiles import StaticFiles
@ -820,13 +820,23 @@ async def proxy_shutdown_event():
async def _initialize_shared_aiohttp_session():
"""Initialize shared aiohttp session for connection reuse with connection limits."""
try:
import socket
from aiohttp import ClientSession, TCPConnector
from litellm.llms.custom_httpx.http_handler import (
_build_aiohttp_keepalive_socket_factory,
)
connector_kwargs: Dict[str, Any] = {
class ConnectorKwargs(TypedDict, total=False):
keepalive_timeout: float
ttl_dns_cache: int
enable_cleanup_closed: bool
limit: int
limit_per_host: int
socket_factory: Callable[[tuple[object, ...]], socket.socket]
connector_kwargs: ConnectorKwargs = {
"keepalive_timeout": AIOHTTP_KEEPALIVE_TIMEOUT,
"ttl_dns_cache": AIOHTTP_TTL_DNS_CACHE,
}
@ -958,7 +968,9 @@ async def proxy_startup_event(app: FastAPI):
async def _run_pw_migration():
try:
result = await migrate_passwords_to_scrypt_async(prisma_client)
result = await migrate_passwords_to_scrypt_async(
cast(PrismaClient, prisma_client) # cast-ok: guarded by the enclosing prisma_client None check
)
verbose_proxy_logger.info(f"Password migration: {result}")
except Exception as e:
verbose_proxy_logger.warning(f"Password migration skipped: {e}")
@ -1141,7 +1153,7 @@ async def proxy_startup_event(app: FastAPI):
await proxy_shutdown_event() # type: ignore[reportGeneralTypeIssues]
def _generate_stable_operation_id(route: Any) -> str:
def _generate_stable_operation_id(route: APIRoute) -> str:
operation_id = re.sub(r"\W", "_", f"{route.name}{route.path_format}")
route_methods = sorted(route.methods or [])
if len(route_methods) == 1:
@ -2782,7 +2794,7 @@ async def update_cache(
Put any alerting logic in here.
"""
values_to_update_in_cache: List[Tuple[Any, Any]] = []
values_to_update_in_cache: list[tuple[str, object]] = []
### UPDATE KEY SPEND ###
async def _update_key_cache(token: str, response_cost: float):
@ -3513,7 +3525,7 @@ def _scrub_guardrail_inner(inner: Dict[str, Any]) -> None:
inner["guardrail"] = None
def _scrub_db_overlay_remote_module_loads(section: str, db_value: Any) -> Any:
def _scrub_db_overlay_remote_module_loads(section: str, db_value: object) -> object:
"""Strip ``s3://`` / ``gcs://`` entries from the DB-overlay value for
fields whose contents reach ``get_instance_fn``. The same scheme is
allowed from a YAML config (the documented operator flow) but a
@ -3915,7 +3927,9 @@ class ProxyConfig:
# if using - db for config - models are in ModelTable
# Make a copy to avoid mutating the original config
config_to_save = new_config.copy()
config_to_save: dict[str, Any] = (
new_config.copy()
) # mutable-ok: keys are popped and re-encrypted in place before the DB write
# environment_variables are persisted to the DB only when a caller
# explicitly opts in. Most callers reach save_config after
@ -5390,8 +5404,12 @@ class ProxyConfig:
added_models += 1
return added_models
def decrypt_model_list_from_db(self, new_models: list) -> list:
_model_list: list = []
def decrypt_model_list_from_db(
self, new_models: Sequence[Any]
) -> list[
dict[str, Any]
]: # mutable-ok: result is handed to litellm.Router(model_list=...), which requires a list of dicts
_model_list: list[dict[str, Any]] = [] # mutable-ok: appended to once per decrypted deployment below
for m in new_models:
_litellm_params = m.litellm_params
if isinstance(_litellm_params, BaseModel):
@ -5400,7 +5418,9 @@ class ProxyConfig:
# decrypt values
for k, v in _litellm_params.items():
_litellm_params[k] = self._resolve_db_litellm_param(key=k, value=v)
_litellm_params = LiteLLM_Params(**_litellm_params)
_litellm_params = cast( # cast-ok: db params dict has untyped values, validated by pydantic at runtime
Callable[..., LiteLLM_Params], LiteLLM_Params
)(**_litellm_params)
else:
verbose_proxy_logger.error(
f"Invalid model added to proxy db. Invalid litellm params. litellm_params={_litellm_params}"
@ -5448,7 +5468,7 @@ class ProxyConfig:
)
return
models_list: list = new_models if isinstance(new_models, list) else []
models_list: list[Any] = new_models if isinstance(new_models, list) else []
if llm_router is None and master_key is not None:
verbose_proxy_logger.debug(f"len new_models: {len(models_list)}")
@ -5620,7 +5640,7 @@ class ProxyConfig:
)
@staticmethod
def _parse_router_settings_value(value: Any) -> Optional[dict]:
def _parse_router_settings_value(value: object) -> dict | None:
"""
Parse a router_settings value that may be a dict or a JSON/YAML string.
@ -6251,7 +6271,9 @@ class ProxyConfig:
if isinstance(litellm_settings, str):
litellm_settings = json.loads(litellm_settings)
mcp_semantic_filter_config = litellm_settings.get("mcp_semantic_tool_filter", None)
mcp_semantic_filter_config = cast( # cast-ok: litellm_settings config rows store a JSON object
dict[str, Any], litellm_settings
).get("mcp_semantic_tool_filter", None)
if mcp_semantic_filter_config is None:
return
@ -6386,7 +6408,9 @@ class ProxyConfig:
if config_record is None or config_record.param_value is None:
return # No configuration found, skip reload
config = config_record.param_value
config = cast( # cast-ok: reload config rows store a JSON object
dict[str, Any], config_record.param_value
)
interval_hours = config.get("interval_hours")
force_reload = config.get("force_reload", False)
@ -6486,7 +6510,9 @@ class ProxyConfig:
if config_record is None or config_record.param_value is None:
return # No configuration found, skip reload
config = config_record.param_value
config = cast( # cast-ok: reload config rows store a JSON object
dict[str, Any], config_record.param_value
)
interval_hours = config.get("interval_hours")
force_reload = config.get("force_reload", False)
@ -7719,7 +7745,11 @@ class ProxyStartupEvent:
"""Initialize MCP semantic tool filter if configured"""
from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook
mcp_semantic_filter_config = litellm_settings.get("mcp_semantic_tool_filter", None)
mcp_semantic_filter_config: dict[str, Any] | None = (
litellm_settings.get( # mutable-ok: forwarded to SemanticToolFilterHook.initialize_from_config, which takes Dict
"mcp_semantic_tool_filter", None
)
)
# Only proceed if the feature is configured and enabled
if not mcp_semantic_filter_config or not mcp_semantic_filter_config.get("enabled", False):
@ -8857,7 +8887,7 @@ async def model_info(
)
def _blocked_response_usage(original_response: Optional[Any]) -> "litellm.Usage":
def _blocked_response_usage(original_response: object | None) -> "litellm.Usage":
"""
Token usage for a synthetic guardrail-blocked response.
@ -8971,7 +9001,9 @@ async def chat_completion(
# Guardrail flagged content in passthrough mode - return 200 with violation message
_data = e.request_data
# Capture logging_obj before post_call_failure_hook pops it from _data.
_logging_obj = _data.get("litellm_logging_obj")
_logging_obj = cast( # cast-ok: chat request data always carries the litellm logging object
LiteLLMLoggingObj, _data.get("litellm_logging_obj")
)
await proxy_logging_obj.post_call_failure_hook(
user_api_key_dict=user_api_key_dict,
original_exception=e,
@ -9022,7 +9054,9 @@ async def chat_completion(
completion_stream=_iterator,
model=data.get("model", ""),
custom_llm_provider="cached_response",
logging_obj=_data.get("litellm_logging_obj", None),
logging_obj=cast( # cast-ok: chat request data always carries the litellm logging object
LiteLLMLoggingObj, _data.get("litellm_logging_obj", None)
),
)
selected_data_generator = select_data_generator(
response=_streaming_response,
@ -11178,13 +11212,20 @@ async def get_all_team_models(
team_db_objects_typed: List[LiteLLM_TeamTable] = []
team_table_from_row = cast( # cast-ok: prisma row dump has untyped values, validated by pydantic at runtime
Callable[..., LiteLLM_TeamTable], LiteLLM_TeamTable
)
if user_teams == "*":
team_db_objects = await TeamRepository(prisma_client).table.find_many()
team_db_objects_typed = [LiteLLM_TeamTable(**team_db_object.model_dump()) for team_db_object in team_db_objects]
team_db_objects_typed = [
team_table_from_row(**team_db_object.model_dump()) for team_db_object in team_db_objects
]
else:
team_db_objects = await TeamRepository(prisma_client).table.find_many(where={"team_id": {"in": user_teams}})
team_db_objects_typed = [LiteLLM_TeamTable(**team_db_object.model_dump()) for team_db_object in team_db_objects]
team_db_objects_typed = [
team_table_from_row(**team_db_object.model_dump()) for team_db_object in team_db_objects
]
team_models = _add_team_models_to_all_models(
team_db_objects_typed=team_db_objects_typed,
@ -11258,7 +11299,9 @@ async def _populate_team_access_on_models(
where={"user_id": user_api_key_dict.user_id}
)
if user_db_object is not None:
user_object = LiteLLM_UserTable(**user_db_object.model_dump())
user_object = cast( # cast-ok: prisma row dump has untyped values, validated by pydantic at runtime
Callable[..., LiteLLM_UserTable], LiteLLM_UserTable
)(**user_db_object.model_dump())
user_teams = user_object.teams or []
direct_access_models = get_direct_access_models(
user_db_object=user_object,
@ -11327,7 +11370,9 @@ def _enrich_model_info_with_litellm_data(
Enriched model dictionary with sensitive info removed
"""
# provided model_info in config.yaml
model_info = model.get("model_info", {})
model_info: dict[str, Any] = model.get(
"model_info", {}
) # mutable-ok: litellm model-cost keys are written into it before it is stored back on the model
if debug is True:
_openai_client = "None"
if llm_router is not None:
@ -11344,7 +11389,7 @@ def _enrich_model_info_with_litellm_data(
# 2nd pass on the model, try seeing if we can find model in litellm model_cost map
if litellm_model_info == {}:
# use litellm_param model_name to get model_info
litellm_params = model.get("litellm_params", {})
litellm_params: dict[str, Any] = model.get("litellm_params", {})
litellm_model = litellm_params.get("model", None)
try:
litellm_model_info = litellm.get_model_info(model=litellm_model)
@ -11375,7 +11420,7 @@ def _enrich_model_info_with_litellm_data(
async def _get_caller_byok_team_scope(
user_api_key_dict: Optional[UserAPIKeyAuth],
prisma_client: Optional[Any],
prisma_client: PrismaClient | None,
) -> Optional[Set[str]]:
"""
Return the team IDs whose BYOK rows the caller is allowed to see via
@ -11432,8 +11477,8 @@ _SORTED_SEARCH_DB_FETCH_CAP = 500
async def _fetch_db_models_for_search(
prisma_client: Any,
proxy_config: Any,
prisma_client: PrismaClient,
proxy_config: ProxyConfig,
search_lower: str,
db_model_ids_in_router: Set[str],
router_models_count: int,
@ -11498,8 +11543,8 @@ async def _fetch_db_models_for_search(
async def _apply_search_filter_to_models(
all_models: List[Dict[str, Any]],
search: str,
prisma_client: Optional[Any],
proxy_config: Any,
prisma_client: PrismaClient | None,
proxy_config: ProxyConfig,
user_api_key_dict: Optional[UserAPIKeyAuth] = None,
page: int = 1,
size: int = 50,
@ -11566,7 +11611,7 @@ async def _apply_search_filter_to_models(
db_model_ids_in_router = set()
for m in filtered_router_models:
model_info = m.get("model_info", {})
model_info: Mapping[str, Any] = m.get("model_info", {})
is_db_model = model_info.get("db_model", False)
model_id = model_info.get("id")
@ -11604,7 +11649,7 @@ async def _apply_search_filter_to_models(
return filtered_router_models + db_models, search_total_count
def _normalize_datetime_for_sorting(dt: Any) -> Optional[datetime]:
def _normalize_datetime_for_sorting(dt: object) -> datetime | None:
"""
Normalize a datetime value to a timezone-aware UTC datetime for sorting.
@ -11674,7 +11719,7 @@ def _sort_models(
reverse = sort_order.lower() == "desc"
def get_sort_key(model: Dict[str, Any]) -> Any:
model_info = model.get("model_info", {})
model_info: Mapping[str, Any] = model.get("model_info", {})
if sort_by == "model_name":
# Team BYOK models persist an internal `model_name` (e.g.
@ -11775,7 +11820,7 @@ def _paginate_models_response(
}
def _team_models_resolve_to_names(team_models: List[str], access_groups: Dict[str, Any]) -> List[str]:
def _team_models_resolve_to_names(team_models: list[str], access_groups: Mapping[str, Sequence[str]]) -> list[str]:
"""Expand team model entries (including access group names) to concrete model names."""
resolved: List[str] = []
for name in team_models:
@ -11793,7 +11838,9 @@ async def _load_team_object_for_model_filter(team_id: str, prisma_client: Prisma
if team_db_object is None:
verbose_proxy_logger.warning(f"Team {team_id} not found in database")
return None
return LiteLLM_TeamTable(**team_db_object.model_dump())
return cast( # cast-ok: prisma row dump has untyped values, validated by pydantic at runtime
Callable[..., LiteLLM_TeamTable], LiteLLM_TeamTable
)(**team_db_object.model_dump())
except Exception as e:
verbose_proxy_logger.exception(f"Error fetching team {team_id}: {str(e)}")
return None
@ -11935,7 +11982,7 @@ async def _filter_models_by_team_id(
# every public model the admin can call.
filtered_models = []
for _model in all_models:
model_info = _model.get("model_info", {})
model_info: dict[str, Any] = _model.get("model_info", {})
model_id = model_info.get("id", None)
# BYOK rows owned by this team are always accessible to it, even if
@ -12286,7 +12333,9 @@ async def model_streaming_metrics(
"""
_all_api_bases = set()
db_response = await prisma_client.db.query_raw(sql_query, _selected_model_group, startTime, endTime)
db_response: Sequence[Mapping[str, Any]] | None = await prisma_client.db.query_raw(
sql_query, _selected_model_group, startTime, endTime
)
_daily_entries: dict = {} # {"Jun 23": {"model1": 0.002, "model2": 0.003}}
if db_response is not None:
for model_data in db_response:
@ -12408,7 +12457,7 @@ async def model_metrics(
avg_latency_per_token DESC;
"""
_all_api_bases = set()
db_response = await prisma_client.db.query_raw(
db_response: Sequence[Mapping[str, Any]] | None = await prisma_client.db.query_raw(
sql_query, _selected_model_group, startTime, endTime, api_key, customer
)
_daily_entries: dict = {} # {"Jun 23": {"model1": 0.002, "model2": 0.003}}
@ -12526,7 +12575,9 @@ ORDER BY
slow_count DESC;
"""
db_response = await prisma_client.db.query_raw(
db_response: Optional[
list[dict[str, Any]]
] = await prisma_client.db.query_raw( # mutable-ok: each row's api_base is rewritten in place below
sql_query,
alerting_threshold,
_selected_model_group,
@ -12971,7 +13022,11 @@ def _get_model_group_info(
_model_group_info = llm_router.get_model_group_info(model_group=model)
if _model_group_info is not None:
model_groups.append(ModelGroupInfoProxy(**_model_group_info.model_dump()))
model_groups.append(
cast( # cast-ok: model group info dump has untyped values, validated by pydantic at runtime
Callable[..., ModelGroupInfoProxy], ModelGroupInfoProxy
)(**_model_group_info.model_dump())
)
else:
model_group_info = ModelGroupInfoProxy(
model_group=model,
@ -13553,7 +13608,7 @@ async def login_v2(request: Request):
from litellm.proxy.utils import get_custom_url
try:
body = await request.json()
body: Mapping[str, Any] = await request.json()
username = str(body.get("username"))
password = str(body.get("password"))
@ -13626,7 +13681,7 @@ async def login_v3(request: Request):
code=status.HTTP_404_NOT_FOUND,
)
body = await request.json()
body: Mapping[str, Any] = await request.json()
username = str(body.get("username"))
password = str(body.get("password"))
@ -13697,7 +13752,7 @@ async def login_v3_exchange(request: Request):
code=status.HTTP_404_NOT_FOUND,
)
body = await request.json()
body: Mapping[str, Any] = await request.json()
code = body.get("code")
if not code:
raise ProxyException(
@ -14690,7 +14745,9 @@ async def update_config_general_settings(
)
try:
ConfigGeneralSettings(**{data.field_name: data.field_value})
cast( # cast-ok: field name/value are dynamic user input, validated by pydantic at runtime
Callable[..., ConfigGeneralSettings], ConfigGeneralSettings
)(**{data.field_name: data.field_value})
except Exception:
raise HTTPException(
status_code=400,
@ -15003,7 +15060,7 @@ def _general_settings_ui_litellm_default(
return False if spec["type"] == "Boolean" else None
def _validate_general_settings_ui_litellm_value(field_name: str, value: Any) -> GeneralSettingsUILiteLLMValue:
def _validate_general_settings_ui_litellm_value(field_name: str, value: object) -> GeneralSettingsUILiteLLMValue:
spec = _GENERAL_SETTINGS_UI_LITELLM_FIELDS[field_name]
field_type = spec["type"]
if value is None or value == "":
@ -15043,7 +15100,7 @@ def _validate_general_settings_ui_litellm_value(field_name: str, value: Any) ->
async def _persist_general_settings_ui_litellm_field(
field_name: str, value: Any, user_api_key_dict: UserAPIKeyAuth
field_name: str, value: object, user_api_key_dict: UserAPIKeyAuth
) -> dict:
validated = _validate_general_settings_ui_litellm_value(field_name, value)
config = await proxy_config.get_config()

File diff suppressed because it is too large Load diff