mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
76b0b10908
commit
6afef76268
13 changed files with 3259 additions and 1217 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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})
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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
Loading…
Add table
Reference in a new issue