Merge pull request #27024 from BerriAI/litellm_yj_may1

[Infra] Merge dev branch
This commit is contained in:
yuneng-jiang 2026-05-01 16:36:24 -07:00 • committed by GitHub
commit 57dd3891fb
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
39 changed files with 6991 additions and 396 deletions

View file

@ -90,6 +90,29 @@ def _extract_cache_read_input_tokens(usage_obj) -> int:
return cache_read_input_tokens
def resolve_langfuse_credentials(
langfuse_public_key=None,
langfuse_secret=None,
langfuse_secret_key=None,
langfuse_host=None,
allow_env_credentials: bool = True,
):
if allow_env_credentials is False and langfuse_host is not None:
secret_key = langfuse_secret or langfuse_secret_key
public_key = langfuse_public_key
else:
secret_key = (
langfuse_secret or langfuse_secret_key or os.getenv("LANGFUSE_SECRET_KEY")
)
public_key = langfuse_public_key or os.getenv("LANGFUSE_PUBLIC_KEY")
resolved_host = langfuse_host or os.getenv(
"LANGFUSE_HOST", "https://cloud.langfuse.com"
)
return public_key, secret_key, resolved_host
class LangFuseLogger:
# Class variables or attributes
def __init__(
@ -98,6 +121,7 @@ class LangFuseLogger:
langfuse_secret=None,
langfuse_host=None,
flush_interval=1,
allow_env_credentials: bool = True,
):
try:
import langfuse
@ -106,11 +130,13 @@ class LangFuseLogger:
raise Exception(
f"\033[91mLangfuse not installed, try running 'pip install langfuse' to fix this error: {e}\n{traceback.format_exc()}\033[0m"
)
# Instance variables
self.secret_key = langfuse_secret or os.getenv("LANGFUSE_SECRET_KEY")
self.public_key = langfuse_public_key or os.getenv("LANGFUSE_PUBLIC_KEY")
self.langfuse_host = langfuse_host or os.getenv(
"LANGFUSE_HOST", "https://cloud.langfuse.com"
self.public_key, self.secret_key, self.langfuse_host = (
resolve_langfuse_credentials(
langfuse_public_key=langfuse_public_key,
langfuse_secret=langfuse_secret,
langfuse_host=langfuse_host,
allow_env_credentials=allow_env_credentials,
)
)
if not (
self.langfuse_host.startswith("http://")
@ -160,9 +186,10 @@ class LangFuseLogger:
project_id = None
if os.getenv("UPSTREAM_LANGFUSE_SECRET_KEY") is not None:
upstream_langfuse_debug_env = os.getenv("UPSTREAM_LANGFUSE_DEBUG")
upstream_langfuse_debug = (
str_to_bool(self.upstream_langfuse_debug)
if self.upstream_langfuse_debug is not None
str_to_bool(upstream_langfuse_debug_env)
if upstream_langfuse_debug_env is not None
else None
)
self.upstream_langfuse_secret_key = os.getenv(
@ -173,7 +200,7 @@ class LangFuseLogger:
)
self.upstream_langfuse_host = os.getenv("UPSTREAM_LANGFUSE_HOST")
self.upstream_langfuse_release = os.getenv("UPSTREAM_LANGFUSE_RELEASE")
self.upstream_langfuse_debug = os.getenv("UPSTREAM_LANGFUSE_DEBUG")
self.upstream_langfuse_debug = upstream_langfuse_debug_env
self.upstream_langfuse = Langfuse(
public_key=self.upstream_langfuse_public_key,
secret_key=self.upstream_langfuse_secret_key,

View file

@ -115,8 +115,10 @@ class LangFuseHandler:
langfuse_logger = LangFuseLogger(
langfuse_public_key=credentials.get("langfuse_public_key"),
langfuse_secret=credentials.get("langfuse_secret"),
langfuse_secret=credentials.get("langfuse_secret")
or credentials.get("langfuse_secret_key"),
langfuse_host=credentials.get("langfuse_host"),
allow_env_credentials=credentials.get("langfuse_host") is None,
)
in_memory_dynamic_logger_cache.set_cache(
credentials=credentials,

View file

@ -20,7 +20,7 @@ from ...litellm_core_utils.specialty_caches.dynamic_logging_cache import (
DynamicLoggingCache,
)
from ..prompt_management_base import PromptManagementBase
from .langfuse import LangFuseLogger
from .langfuse import LangFuseLogger, resolve_langfuse_credentials
from .langfuse_handler import LangFuseHandler
if TYPE_CHECKING:
@ -46,6 +46,7 @@ def langfuse_client_init(
langfuse_secret_key=None,
langfuse_host=None,
flush_interval=1,
allow_env_credentials: bool = True,
) -> LangfuseClass:
"""
Initialize Langfuse client with caching to prevent multiple initializations.
@ -70,14 +71,12 @@ def langfuse_client_init(
f"\033[91mLangfuse not installed, try running 'pip install langfuse' to fix this error: {e}\n\033[0m"
)
# Instance variables
secret_key = (
langfuse_secret or langfuse_secret_key or os.getenv("LANGFUSE_SECRET_KEY")
)
public_key = langfuse_public_key or os.getenv("LANGFUSE_PUBLIC_KEY")
langfuse_host = langfuse_host or os.getenv(
"LANGFUSE_HOST", "https://cloud.langfuse.com"
public_key, secret_key, langfuse_host = resolve_langfuse_credentials(
langfuse_public_key=langfuse_public_key,
langfuse_secret=langfuse_secret,
langfuse_secret_key=langfuse_secret_key,
langfuse_host=langfuse_host,
allow_env_credentials=allow_env_credentials,
)
if not (
@ -222,6 +221,7 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge
langfuse_secret=dynamic_callback_params.get("langfuse_secret"),
langfuse_secret_key=dynamic_callback_params.get("langfuse_secret_key"),
langfuse_host=dynamic_callback_params.get("langfuse_host"),
allow_env_credentials=dynamic_callback_params.get("langfuse_host") is None,
)
langfuse_prompt_client = self._get_prompt_from_id(
langfuse_prompt_id=prompt_id,
@ -246,6 +246,7 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge
langfuse_secret=dynamic_callback_params.get("langfuse_secret"),
langfuse_secret_key=dynamic_callback_params.get("langfuse_secret_key"),
langfuse_host=dynamic_callback_params.get("langfuse_host"),
allow_env_credentials=dynamic_callback_params.get("langfuse_host") is None,
)
langfuse_prompt_client = self._get_prompt_from_id(
langfuse_prompt_id=prompt_id,

View file

@ -112,17 +112,28 @@ class LangsmithLogger(CustomBatchLogger):
langsmith_project: Optional[str] = None,
langsmith_base_url: Optional[str] = None,
langsmith_tenant_id: Optional[str] = None,
allow_env_credentials: bool = True,
) -> LangsmithCredentialsObject:
_credentials_api_key = langsmith_api_key or os.getenv("LANGSMITH_API_KEY")
_credentials_project = (
langsmith_project or os.getenv("LANGSMITH_PROJECT") or "litellm-completion"
)
_credentials_base_url = (
langsmith_base_url
or os.getenv("LANGSMITH_BASE_URL")
or "https://api.smith.langchain.com"
)
_credentials_tenant_id = langsmith_tenant_id or os.getenv("LANGSMITH_TENANT_ID")
if allow_env_credentials is False and langsmith_base_url is not None:
_credentials_api_key = langsmith_api_key
_credentials_project = langsmith_project or "litellm-completion"
_credentials_base_url = langsmith_base_url
_credentials_tenant_id = langsmith_tenant_id
else:
_credentials_api_key = langsmith_api_key or os.getenv("LANGSMITH_API_KEY")
_credentials_project = (
langsmith_project
or os.getenv("LANGSMITH_PROJECT")
or "litellm-completion"
)
_credentials_base_url = (
langsmith_base_url
or os.getenv("LANGSMITH_BASE_URL")
or "https://api.smith.langchain.com"
)
_credentials_tenant_id = langsmith_tenant_id or os.getenv(
"LANGSMITH_TENANT_ID"
)
return LangsmithCredentialsObject(
LANGSMITH_API_KEY=_credentials_api_key,
@ -540,6 +551,10 @@ class LangsmithLogger(CustomBatchLogger):
langsmith_tenant_id=standard_callback_dynamic_params.get(
"langsmith_tenant_id", None
),
allow_env_credentials=standard_callback_dynamic_params.get(
"langsmith_base_url", None
)
is None,
)
else:
credentials = self.default_credentials

View file

@ -3242,10 +3242,15 @@ class Logging(LiteLLMLoggingBaseClass):
),
langfuse_secret=self.standard_callback_dynamic_params.get(
"langfuse_secret"
),
)
or self.standard_callback_dynamic_params.get("langfuse_secret_key"),
langfuse_host=self.standard_callback_dynamic_params.get(
"langfuse_host"
),
allow_env_credentials=self.standard_callback_dynamic_params.get(
"langfuse_host"
)
is None,
)
return langFuseLogger
@ -4720,7 +4725,7 @@ class StandardLoggingPayloadSetup:
):
for key, value in litellm_params["metadata"].items():
# Skip non-serializable objects like UserAPIKeyAuth
if key == "user_api_key_auth":
if key in {"user_api_key_auth", "user_api_key_budget_reservation"}:
continue
merged_metadata[key] = value

View file

@ -3616,7 +3616,7 @@
},
"get": {
"description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)",
"operationId": "anthropic_proxy_route_anthropic__endpoint__get",
"operationId": "anthropic_proxy_route_anthropic__endpoint__delete",
"parameters": [
{
"in": "path",
@ -3660,7 +3660,7 @@
},
"patch": {
"description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)",
"operationId": "anthropic_proxy_route_anthropic__endpoint__patch",
"operationId": "anthropic_proxy_route_anthropic__endpoint__delete",
"parameters": [
{
"in": "path",
@ -3704,7 +3704,7 @@
},
"post": {
"description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)",
"operationId": "anthropic_proxy_route_anthropic__endpoint__post",
"operationId": "anthropic_proxy_route_anthropic__endpoint__delete",
"parameters": [
{
"in": "path",
@ -3748,7 +3748,7 @@
},
"put": {
"description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)",
"operationId": "anthropic_proxy_route_anthropic__endpoint__put",
"operationId": "anthropic_proxy_route_anthropic__endpoint__delete",
"parameters": [
{
"in": "path",
@ -13299,7 +13299,7 @@
},
"get": {
"description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)",
"operationId": "langfuse_proxy_route_langfuse__endpoint__get",
"operationId": "langfuse_proxy_route_langfuse__endpoint__delete",
"parameters": [
{
"in": "path",
@ -13338,7 +13338,7 @@
},
"patch": {
"description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)",
"operationId": "langfuse_proxy_route_langfuse__endpoint__patch",
"operationId": "langfuse_proxy_route_langfuse__endpoint__delete",
"parameters": [
{
"in": "path",
@ -13377,7 +13377,7 @@
},
"post": {
"description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)",
"operationId": "langfuse_proxy_route_langfuse__endpoint__post",
"operationId": "langfuse_proxy_route_langfuse__endpoint__delete",
"parameters": [
{
"in": "path",
@ -13416,7 +13416,7 @@
},
"put": {
"description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)",
"operationId": "langfuse_proxy_route_langfuse__endpoint__put",
"operationId": "langfuse_proxy_route_langfuse__endpoint__delete",
"parameters": [
{
"in": "path",
@ -14008,7 +14008,7 @@
"/mcp-rest/test/connection": {
"post": {
"description": "Test if we can connect to the provided MCP server before adding it",
"operationId": "test_connection_mcp_rest_test_connection_post",
"operationId": "test_connection_mcp_rest_test_connection_post_2",
"requestBody": {
"content": {
"application/json": {
@ -14053,7 +14053,7 @@
"/mcp-rest/test/tools/list": {
"post": {
"description": "Preview tools available from MCP server before adding it",
"operationId": "test_tools_list_mcp_rest_test_tools_list_post",
"operationId": "test_tools_list_mcp_rest_test_tools_list_post_2",
"requestBody": {
"content": {
"application/json": {
@ -14098,7 +14098,7 @@
"/mcp-rest/tools/call": {
"post": {
"description": "REST API to call a specific MCP tool with the provided arguments",
"operationId": "call_tool_rest_api_mcp_rest_tools_call_post",
"operationId": "call_tool_rest_api_mcp_rest_tools_call_post_2",
"responses": {
"200": {
"content": {
@ -14123,7 +14123,7 @@
"/mcp-rest/tools/list": {
"get": {
"description": "List all available tools with information about the server they belong to.\n\nExample response:\n{\n \"tools\": [\n {\n \"name\": \"create_zap\",\n \"description\": \"Create a new zap\",\n \"inputSchema\": \"tool_input_schema\",\n \"mcp_info\": {\n \"server_name\": \"zapier\",\n \"logo_url\": \"https://www.zapier.com/logo.png\",\n }\n }\n ],\n \"error\": null,\n \"message\": \"Successfully retrieved tools\"\n}",
"operationId": "list_tool_rest_api_mcp_rest_tools_list_get",
"operationId": "list_tool_rest_api_mcp_rest_tools_list_get_2",
"parameters": [
{
"description": "The server id to list tools for",
@ -21896,7 +21896,7 @@
"/policies/usage/overview": {
"get": {
"description": "Return policy performance overview for the dashboard.",
"operationId": "policies_usage_overview_policies_usage_overview_get",
"operationId": "policies_usage_overview_policies_usage_overview_get_2",
"parameters": [
{
"description": "YYYY-MM-DD",
@ -22521,7 +22521,7 @@
"/policies/attachments/estimate-impact": {
"post": {
"description": "Estimate how many keys and teams would be affected by a policy attachment.\n\nUse this before creating an attachment to preview the blast radius.\n\nExample Request:\n```bash\ncurl -X POST \"http://localhost:4000/policies/attachments/estimate-impact\" \\\n -H \"Authorization: Bearer <your_api_key>\" \\\n -H \"Content-Type: application/json\" \\\n -d '{\n \"policy_name\": \"hipaa-compliance\",\n \"tags\": [\"healthcare\", \"health-*\"]\n }'\n```",
"operationId": "estimate_attachment_impact_policies_attachments_estimate_impact_post",
"operationId": "estimate_attachment_impact_policies_attachments_estimate_impact_post_2",
"requestBody": {
"content": {
"application/json": {
@ -22568,7 +22568,7 @@
"/policies/resolve": {
"post": {
"description": "Resolve which policies and guardrails apply for a given context.\n\nUse this endpoint to debug \"what guardrails would apply to a request\nwith this team/key/model/tags combination?\"\n\nExample Request:\n```bash\ncurl -X POST \"http://localhost:4000/policies/resolve\" \\\n -H \"Authorization: Bearer <your_api_key>\" \\\n -H \"Content-Type: application/json\" \\\n -d '{\n \"tags\": [\"healthcare\"],\n \"model\": \"gpt-4\"\n }'\n```",
"operationId": "resolve_policies_for_context_policies_resolve_post",
"operationId": "resolve_policies_for_context_policies_resolve_post_2",
"parameters": [
{
"description": "Force a DB sync before resolving. Default uses in-memory cache.",
@ -26922,7 +26922,7 @@
},
"get": {
"description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset/<name>/mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.",
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_get",
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_delete",
"parameters": [
{
"in": "path",
@ -26961,7 +26961,7 @@
},
"head": {
"description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset/<name>/mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.",
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_head",
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_delete",
"parameters": [
{
"in": "path",
@ -27000,7 +27000,7 @@
},
"options": {
"description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset/<name>/mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.",
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_options",
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_delete",
"parameters": [
{
"in": "path",
@ -27039,7 +27039,7 @@
},
"patch": {
"description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset/<name>/mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.",
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_patch",
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_delete",
"parameters": [
{
"in": "path",
@ -27078,7 +27078,7 @@
},
"post": {
"description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset/<name>/mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.",
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_post",
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_delete",
"parameters": [
{
"in": "path",
@ -27117,7 +27117,7 @@
},
"put": {
"description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset/<name>/mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.",
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put",
"operationId": "toolset_mcp_route_toolset__toolset_name__mcp_delete",
"parameters": [
{
"in": "path",
@ -28329,7 +28329,7 @@
"/v1/vector_stores": {
"get": {
"description": "List vector stores.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/list",
"operationId": "vector_store_list_v1_vector_stores_get",
"operationId": "vector_store_list_v1_vector_stores_get_2",
"parameters": [
{
"in": "query",
@ -28430,7 +28430,7 @@
},
"post": {
"description": "Create a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/create\n\nSupports target_model_names parameter for creating vector stores across multiple models:\n```json\n{\n \"name\": \"my-vector-store\",\n \"target_model_names\": \"gpt-4,gemini-2.0\"\n}\n```",
"operationId": "vector_store_create_v1_vector_stores_post",
"operationId": "vector_store_create_v1_vector_stores_post_2",
"responses": {
"200": {
"content": {
@ -28455,7 +28455,7 @@
"/v1/vector_stores/{vector_store_id}": {
"delete": {
"description": "Delete a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/delete",
"operationId": "vector_store_delete_v1_vector_stores__vector_store_id__delete",
"operationId": "vector_store_delete_v1_vector_stores__vector_store_id__delete_2",
"parameters": [
{
"in": "path",
@ -28499,7 +28499,7 @@
},
"get": {
"description": "Retrieve a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/retrieve",
"operationId": "vector_store_retrieve_v1_vector_stores__vector_store_id__get",
"operationId": "vector_store_retrieve_v1_vector_stores__vector_store_id__get_2",
"parameters": [
{
"in": "path",
@ -28543,7 +28543,7 @@
},
"post": {
"description": "Update a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/modify",
"operationId": "vector_store_update_v1_vector_stores__vector_store_id__post",
"operationId": "vector_store_update_v1_vector_stores__vector_store_id__post_2",
"parameters": [
{
"in": "path",
@ -28588,7 +28588,7 @@
},
"/v1/vector_stores/{vector_store_id}/files": {
"get": {
"operationId": "vector_store_file_list_v1_vector_stores__vector_store_id__files_get",
"operationId": "vector_store_file_list_v1_vector_stores__vector_store_id__files_get_2",
"parameters": [
{
"in": "path",
@ -28631,7 +28631,7 @@
]
},
"post": {
"operationId": "vector_store_file_create_v1_vector_stores__vector_store_id__files_post",
"operationId": "vector_store_file_create_v1_vector_stores__vector_store_id__files_post_2",
"parameters": [
{
"in": "path",
@ -28676,7 +28676,7 @@
},
"/v1/vector_stores/{vector_store_id}/files/{file_id}": {
"delete": {
"operationId": "vector_store_file_delete_v1_vector_stores__vector_store_id__files__file_id__delete",
"operationId": "vector_store_file_delete_v1_vector_stores__vector_store_id__files__file_id__delete_2",
"parameters": [
{
"in": "path",
@ -28728,7 +28728,7 @@
]
},
"get": {
"operationId": "vector_store_file_retrieve_v1_vector_stores__vector_store_id__files__file_id__get",
"operationId": "vector_store_file_retrieve_v1_vector_stores__vector_store_id__files__file_id__get_2",
"parameters": [
{
"in": "path",
@ -28780,7 +28780,7 @@
]
},
"post": {
"operationId": "vector_store_file_update_v1_vector_stores__vector_store_id__files__file_id__post",
"operationId": "vector_store_file_update_v1_vector_stores__vector_store_id__files__file_id__post_2",
"parameters": [
{
"in": "path",
@ -28834,7 +28834,7 @@
},
"/v1/vector_stores/{vector_store_id}/files/{file_id}/content": {
"get": {
"operationId": "vector_store_file_content_v1_vector_stores__vector_store_id__files__file_id__content_get",
"operationId": "vector_store_file_content_v1_vector_stores__vector_store_id__files__file_id__content_get_2",
"parameters": [
{
"in": "path",
@ -28889,7 +28889,7 @@
"/v1/vector_stores/{vector_store_id}/search": {
"post": {
"description": "Search a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/search",
"operationId": "vector_store_search_v1_vector_stores__vector_store_id__search_post",
"operationId": "vector_store_search_v1_vector_stores__vector_store_id__search_post_2",
"parameters": [
{
"in": "path",
@ -28935,7 +28935,7 @@
"/vector_stores": {
"get": {
"description": "List vector stores.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/list",
"operationId": "vector_store_list_vector_stores_get",
"operationId": "vector_store_list_vector_stores_get_2",
"parameters": [
{
"in": "query",
@ -29036,7 +29036,7 @@
},
"post": {
"description": "Create a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/create\n\nSupports target_model_names parameter for creating vector stores across multiple models:\n```json\n{\n \"name\": \"my-vector-store\",\n \"target_model_names\": \"gpt-4,gemini-2.0\"\n}\n```",
"operationId": "vector_store_create_vector_stores_post",
"operationId": "vector_store_create_vector_stores_post_2",
"responses": {
"200": {
"content": {
@ -29061,7 +29061,7 @@
"/vector_stores/{vector_store_id}": {
"delete": {
"description": "Delete a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/delete",
"operationId": "vector_store_delete_vector_stores__vector_store_id__delete",
"operationId": "vector_store_delete_vector_stores__vector_store_id__delete_2",
"parameters": [
{
"in": "path",
@ -29105,7 +29105,7 @@
},
"get": {
"description": "Retrieve a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/retrieve",
"operationId": "vector_store_retrieve_vector_stores__vector_store_id__get",
"operationId": "vector_store_retrieve_vector_stores__vector_store_id__get_2",
"parameters": [
{
"in": "path",
@ -29149,7 +29149,7 @@
},
"post": {
"description": "Update a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/modify",
"operationId": "vector_store_update_vector_stores__vector_store_id__post",
"operationId": "vector_store_update_vector_stores__vector_store_id__post_2",
"parameters": [
{
"in": "path",
@ -29194,7 +29194,7 @@
},
"/vector_stores/{vector_store_id}/files": {
"get": {
"operationId": "vector_store_file_list_vector_stores__vector_store_id__files_get",
"operationId": "vector_store_file_list_vector_stores__vector_store_id__files_get_2",
"parameters": [
{
"in": "path",
@ -29237,7 +29237,7 @@
]
},
"post": {
"operationId": "vector_store_file_create_vector_stores__vector_store_id__files_post",
"operationId": "vector_store_file_create_vector_stores__vector_store_id__files_post_2",
"parameters": [
{
"in": "path",
@ -29282,7 +29282,7 @@
},
"/vector_stores/{vector_store_id}/files/{file_id}": {
"delete": {
"operationId": "vector_store_file_delete_vector_stores__vector_store_id__files__file_id__delete",
"operationId": "vector_store_file_delete_vector_stores__vector_store_id__files__file_id__delete_2",
"parameters": [
{
"in": "path",
@ -29334,7 +29334,7 @@
]
},
"get": {
"operationId": "vector_store_file_retrieve_vector_stores__vector_store_id__files__file_id__get",
"operationId": "vector_store_file_retrieve_vector_stores__vector_store_id__files__file_id__get_2",
"parameters": [
{
"in": "path",
@ -29386,7 +29386,7 @@
]
},
"post": {
"operationId": "vector_store_file_update_vector_stores__vector_store_id__files__file_id__post",
"operationId": "vector_store_file_update_vector_stores__vector_store_id__files__file_id__post_2",
"parameters": [
{
"in": "path",
@ -29440,7 +29440,7 @@
},
"/vector_stores/{vector_store_id}/files/{file_id}/content": {
"get": {
"operationId": "vector_store_file_content_vector_stores__vector_store_id__files__file_id__content_get",
"operationId": "vector_store_file_content_vector_stores__vector_store_id__files__file_id__content_get_2",
"parameters": [
{
"in": "path",
@ -29495,7 +29495,7 @@
"/vector_stores/{vector_store_id}/search": {
"post": {
"description": "Search a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/search",
"operationId": "vector_store_search_vector_stores__vector_store_id__search_post",
"operationId": "vector_store_search_vector_stores__vector_store_id__search_post_2",
"parameters": [
{
"in": "path",

View file

@ -8,12 +8,35 @@ any drift as a neutral check.
"""
import json
import re
import sys
from pathlib import Path
from typing import Dict, Optional
from typing import Dict, Optional, Set
SNAPSHOT_FILE = Path(__file__).parent / "_lazy_openapi_snapshot.json"
HTTP_METHODS = {"delete", "get", "head", "options", "patch", "post", "put"}
HTTP_METHOD_SUFFIXES = {
"delete",
"get",
"head",
"options",
"patch",
"post",
"put",
"trace",
}
def _stabilize_multi_method_route_ids(routes) -> None:
"""FastAPI derives route IDs from a set of methods; make snapshots stable."""
for route in routes:
methods = sorted(getattr(route, "methods", None) or [])
if len(methods) <= 1 or not getattr(route, "path_format", None):
continue
operation_id = f"{route.name}{route.path_format}"
operation_id = re.sub(r"\W", "_", operation_id)
route.unique_id = f"{operation_id}_{methods[0].lower()}"
def load_snapshot() -> Optional[Dict[str, Dict]]:
@ -38,12 +61,12 @@ def _normalize_operation_ids(paths: Dict[str, Dict]) -> None:
if not isinstance(path_ops, dict):
continue
methods = {method for method in path_ops if method in HTTP_METHODS}
methods = {method for method in path_ops if method in HTTP_METHOD_SUFFIXES}
if not methods:
continue
for method, operation in path_ops.items():
if method not in HTTP_METHODS or not isinstance(operation, dict):
if method not in HTTP_METHOD_SUFFIXES or not isinstance(operation, dict):
continue
operation_id = operation.get("operationId")
@ -65,7 +88,7 @@ def generate_snapshot() -> Dict[str, Dict]:
from fastapi.openapi.utils import get_openapi
from litellm.proxy._lazy_features import LAZY_FEATURES
from litellm.proxy.proxy_server import app
from litellm.proxy.proxy_server import app, ensure_unique_openapi_operation_ids
for feat in LAZY_FEATURES:
if feat.module_path in sys.modules:
@ -77,6 +100,7 @@ def generate_snapshot() -> Dict[str, Dict]:
sys.stderr.write(f"warning: skip {feat.name}: {exc}\n")
fragments: Dict[str, Dict] = {}
used_operation_ids: Set[str] = set()
for feat in LAZY_FEATURES:
feat_routes = [
r
@ -85,14 +109,24 @@ def generate_snapshot() -> Dict[str, Dict]:
]
if not feat_routes:
continue
_stabilize_multi_method_route_ids(feat_routes)
full = get_openapi(title=app.title, version=app.version, routes=feat_routes)
paths = full.get("paths", {})
_normalize_operation_ids(paths)
# Group all of a feature's routes under one tag.
for path_ops in paths.values():
for op in path_ops.values():
for path_ops in full.get("paths", {}).values():
for method, op in path_ops.items():
if isinstance(op, dict):
operation_id = op.get("operationId")
if isinstance(operation_id, str):
for suffix in HTTP_METHOD_SUFFIXES:
if operation_id.endswith(f"_{suffix}"):
op["operationId"] = (
operation_id[: -len(suffix)] + method
)
break
op["tags"] = [feat.name]
full = ensure_unique_openapi_operation_ids(full, used_operation_ids)
fragments[feat.name] = {
"paths": paths,
"components": {"schemas": full.get("components", {}).get("schemas", {})},

View file

@ -2579,6 +2579,7 @@ class UserAPIKeyAuth(
user_spend: Optional[float] = None
user_max_budget: Optional[float] = None
request_route: Optional[str] = None
budget_reservation: Optional[Dict[str, Any]] = Field(default=None, exclude=True)
user: Optional[Any] = None # Expanded user object when expand=user is used
created_by_user: Optional[Any] = (
None # Expanded created_by user when expand=user is used

View file

@ -60,6 +60,10 @@ from litellm.proxy._types import (
UserAPIKeyAuth,
)
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.common_utils.http_parsing_utils import (
_safe_get_request_headers,
_safe_get_request_query_params,
)
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
from litellm.proxy.guardrails.tool_name_extraction import (
TOOL_CAPABLE_CALL_TYPES,
@ -486,7 +490,10 @@ async def common_checks( # noqa: PLR0915
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
_model: Optional[Union[str, List[str]]] = get_model_from_request(
request_body, route
request_data=request_body,
route=route,
request_headers=_safe_get_request_headers(request=request),
request_query_params=_safe_get_request_query_params(request=request),
)
# 1. If team is blocked
@ -656,13 +663,7 @@ async def common_checks( # noqa: PLR0915
end_user_object is not None
and end_user_object.litellm_budget_table is not None
):
end_user_budget = end_user_object.litellm_budget_table.max_budget
if end_user_budget is not None and end_user_object.spend > end_user_budget:
raise litellm.BudgetExceededError(
current_cost=end_user_object.spend,
max_budget=end_user_budget,
message=f"ExceededBudget: End User={end_user_object.user_id} over budget. Spend={end_user_object.spend}, Budget={end_user_budget}",
)
await _check_end_user_budget(end_user_obj=end_user_object, route=route)
_enforce_user_param_check(general_settings, request, request_body, route)
_reject_clientside_metadata_tags_check(general_settings, request_body, route)
@ -1012,7 +1013,7 @@ async def _apply_default_budget_to_end_user(
return end_user_obj
def _check_end_user_budget(
async def _check_end_user_budget(
end_user_obj: LiteLLM_EndUserTable,
route: str,
) -> None:
@ -1033,11 +1034,20 @@ def _check_end_user_budget(
return
end_user_budget = end_user_obj.litellm_budget_table.max_budget
if end_user_budget is not None and end_user_obj.spend > end_user_budget:
if end_user_budget is None:
return
from litellm.proxy.proxy_server import get_current_spend
end_user_spend = await get_current_spend(
counter_key=f"spend:end_user:{end_user_obj.user_id}",
fallback_spend=end_user_obj.spend or 0.0,
)
if end_user_spend > end_user_budget:
raise litellm.BudgetExceededError(
current_cost=end_user_obj.spend,
current_cost=end_user_spend,
max_budget=end_user_budget,
message=f"ExceededBudget: End User={end_user_obj.user_id} over budget. Spend={end_user_obj.spend}, Budget={end_user_budget}",
message=f"ExceededBudget: End User={end_user_obj.user_id} over budget. Spend={end_user_spend}, Budget={end_user_budget}",
)
@ -1091,7 +1101,7 @@ async def get_end_user_object(
)
# Check budget limits
_check_end_user_budget(end_user_obj=return_obj, route=route)
await _check_end_user_budget(end_user_obj=return_obj, route=route)
return return_obj
@ -1124,7 +1134,7 @@ async def get_end_user_object(
)
# Check budget limits
_check_end_user_budget(end_user_obj=_response, route=route)
await _check_end_user_budget(end_user_obj=_response, route=route)
return _response
@ -1616,9 +1626,12 @@ async def _cache_key_object(
## CACHE REFRESH TIME
user_api_key_obj.last_refreshed_at = time.time()
cached_key_obj = _copy_user_api_key_auth_for_cache(
user_api_key_obj=user_api_key_obj
)
await _cache_management_object(
key=key,
value=user_api_key_obj,
value=cached_key_obj,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
model_type=UserAPIKeyAuth,
@ -2348,7 +2361,7 @@ async def get_key_object(
model_type=UserAPIKeyAuth,
)
if user_api_key_auth is not None:
return user_api_key_auth
return _copy_user_api_key_auth_for_cache(user_api_key_obj=user_api_key_auth)
if check_cache_only:
raise Exception(
@ -2401,6 +2414,16 @@ async def get_key_object(
return _response
def _copy_user_api_key_auth_for_cache(
user_api_key_obj: UserAPIKeyAuth,
) -> UserAPIKeyAuth:
copied_key_obj = user_api_key_obj.model_copy()
copied_key_obj.budget_reservation = None
copied_key_obj.parent_otel_span = None
copied_key_obj.request_route = None
return copied_key_obj
@log_db_metrics
async def get_object_permission(
object_permission_id: str,
@ -3967,13 +3990,19 @@ async def _tag_max_budget_check(
if (
tag_object.litellm_budget_table is not None
and tag_object.litellm_budget_table.max_budget is not None
and tag_object.spend is not None
and tag_object.spend > tag_object.litellm_budget_table.max_budget
):
from litellm.proxy.proxy_server import get_current_spend
tag_spend = await get_current_spend(
counter_key=f"spend:tag:{tag_name}",
fallback_spend=tag_object.spend or 0.0,
)
if tag_spend <= tag_object.litellm_budget_table.max_budget:
continue
raise litellm.BudgetExceededError(
current_cost=tag_object.spend,
current_cost=tag_spend,
max_budget=tag_object.litellm_budget_table.max_budget,
message=f"Budget has been exceeded! Tag={tag_name} Current cost: {tag_object.spend}, Max budget: {tag_object.litellm_budget_table.max_budget}",
message=f"Budget has been exceeded! Tag={tag_name} Current cost: {tag_spend}, Max budget: {tag_object.litellm_budget_table.max_budget}",
)

View file

@ -2,7 +2,7 @@ import os
import re
import sys
from functools import lru_cache
from typing import Any, List, Optional, Tuple
from typing import Any, Dict, List, Mapping, Optional, Tuple, Union
from fastapi import HTTPException, Request, status
@ -976,20 +976,257 @@ def get_end_user_id_from_request_body(
return None
def get_model_from_request(
request_data: dict, route: str
) -> Optional[Union[str, List[str]]]:
# First try to get model from request_data
model = request_data.get("model") or request_data.get("target_model_names")
MODEL_ROUTING_HEADER_NAME = "x-litellm-model"
_MODEL_ROUTING_ROUTE_MARKERS = (
"/files",
"/batches",
"/vector_stores",
"/skills",
"/evals",
"/fine_tuning",
"/videos",
)
_MODEL_ROUTING_HEADER_OR_QUERY_ROUTE_MARKERS = (
"/files",
"/batches",
"/skills",
"/evals",
)
_MODEL_ROUTING_QUERY_TARGET_MODEL_ROUTE_MARKERS = (
"/files",
"/batches",
"/fine_tuning",
)
_MODEL_ROUTING_BODY_TARGET_MODEL_ROUTE_MARKERS = (
"/files",
"/batches",
"/vector_stores",
)
_MODEL_ROUTING_COMPLETION_MODEL_ROUTE_MARKERS = ("/evals",)
_MODEL_ROUTING_ID_FIELDS = (
"file_id",
"input_file_id",
"output_file_id",
"error_file_id",
"batch_id",
"fine_tuning_job_id",
"training_file",
"validation_file",
"vector_store_id",
"video_id",
"character_id",
)
if model is not None:
model_names = model.split(",")
if len(model_names) == 1:
model = model_names[0].strip()
def _append_model_candidates(candidates: List[str], value: Any) -> None:
if value is None:
return
values = value if isinstance(value, (list, tuple, set)) else [value]
for item in values:
if item is None:
continue
if isinstance(item, str):
model_names = [model.strip() for model in item.split(",")]
else:
model = [m.strip() for m in model_names]
model_names = [str(item).strip()]
candidates.extend(model for model in model_names if model)
# If model not in request_data, try to extract from route
def _dedupe_model_candidates(candidates: List[str]) -> List[str]:
deduped: List[str] = []
for model in candidates:
if model not in deduped:
deduped.append(model)
return deduped
def _get_case_insensitive_mapping_value(
mapping: Optional[Mapping[str, Any]], key: str
) -> Any:
if not mapping:
return None
if key in mapping:
return mapping[key]
key_lower = key.lower()
for mapping_key, value in mapping.items():
if str(mapping_key).lower() == key_lower:
return value
return None
def _route_matches_any_marker(route: str, markers: Tuple[str, ...]) -> bool:
normalized_route = route.lower()
return any(marker in normalized_route for marker in markers)
def _route_uses_model_routing_sources(route: str) -> bool:
return _route_matches_any_marker(route=route, markers=_MODEL_ROUTING_ROUTE_MARKERS)
def _extract_models_from_managed_resource_id(
resource_id: Any, resource_id_field: Optional[str] = None
) -> List[str]:
if not isinstance(resource_id, str) or not resource_id:
return []
candidates: List[str] = []
try:
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
decode_model_from_file_id,
get_model_id_from_unified_batch_id,
get_models_from_unified_file_id,
)
_append_model_candidates(
candidates=candidates, value=decode_model_from_file_id(resource_id)
)
unified_file_id = _is_base64_encoded_unified_file_id(resource_id)
if unified_file_id:
_append_model_candidates(
candidates=candidates,
value=get_models_from_unified_file_id(unified_file_id),
)
_append_model_candidates(
candidates=candidates,
value=get_model_id_from_unified_batch_id(unified_file_id),
)
except Exception as e:
verbose_proxy_logger.debug(
"Unable to extract model from managed file/batch ID: %s", str(e)
)
try:
from litellm.llms.base_llm.managed_resources.utils import parse_unified_id
parsed_id = parse_unified_id(resource_id)
if parsed_id:
_append_model_candidates(
candidates=candidates, value=parsed_id.get("model_id")
)
_append_model_candidates(
candidates=candidates, value=parsed_id.get("target_model_names")
)
except Exception as e:
verbose_proxy_logger.debug(
"Unable to extract model from unified managed resource ID: %s", str(e)
)
if resource_id_field in ("video_id", "character_id"):
try:
from litellm.types.videos.utils import (
decode_character_id_with_provider,
decode_video_id_with_provider,
)
if resource_id_field == "video_id":
_append_model_candidates(
candidates=candidates,
value=decode_video_id_with_provider(resource_id).get("model_id"),
)
else:
_append_model_candidates(
candidates=candidates,
value=decode_character_id_with_provider(resource_id).get(
"model_id"
),
)
except Exception as e:
verbose_proxy_logger.debug(
"Unable to extract model from managed video/character ID: %s", str(e)
)
return _dedupe_model_candidates(candidates)
def _extract_model_candidates_from_request(
request_data: dict,
route: str,
request_headers: Optional[Mapping[str, Any]] = None,
request_query_params: Optional[Mapping[str, Any]] = None,
) -> List[str]:
candidates: List[str] = []
uses_model_routing_sources = _route_uses_model_routing_sources(route=route)
uses_header_or_query_model_sources = _route_matches_any_marker(
route=route, markers=_MODEL_ROUTING_HEADER_OR_QUERY_ROUTE_MARKERS
)
uses_query_target_model_sources = _route_matches_any_marker(
route=route, markers=_MODEL_ROUTING_QUERY_TARGET_MODEL_ROUTE_MARKERS
)
uses_body_target_model_sources = _route_matches_any_marker(
route=route, markers=_MODEL_ROUTING_BODY_TARGET_MODEL_ROUTE_MARKERS
)
uses_completion_model_sources = _route_matches_any_marker(
route=route, markers=_MODEL_ROUTING_COMPLETION_MODEL_ROUTE_MARKERS
)
body_model = request_data.get("model")
_append_model_candidates(candidates, body_model)
if uses_body_target_model_sources or not body_model:
_append_model_candidates(candidates, request_data.get("target_model_names"))
if uses_completion_model_sources and isinstance(
request_data.get("completion"), dict
):
_append_model_candidates(candidates, request_data["completion"].get("model"))
if uses_model_routing_sources:
if uses_header_or_query_model_sources:
_append_model_candidates(
candidates,
_get_case_insensitive_mapping_value(request_query_params, "model"),
)
_append_model_candidates(
candidates,
_get_case_insensitive_mapping_value(
request_headers, MODEL_ROUTING_HEADER_NAME
),
)
if uses_query_target_model_sources:
_append_model_candidates(
candidates,
_get_case_insensitive_mapping_value(
request_query_params, "target_model_names"
),
)
for field in _MODEL_ROUTING_ID_FIELDS:
_append_model_candidates(
candidates,
_extract_models_from_managed_resource_id(
request_data.get(field), resource_id_field=field
),
)
return _dedupe_model_candidates(candidates)
def _format_model_candidates(
candidates: List[str],
) -> Optional[Union[str, List[str]]]:
if not candidates:
return None
if len(candidates) == 1:
return candidates[0]
return candidates
def get_model_from_request(
request_data: dict,
route: str,
request_headers: Optional[Mapping[str, Any]] = None,
request_query_params: Optional[Mapping[str, Any]] = None,
) -> Optional[Union[str, List[str]]]:
candidates = _extract_model_candidates_from_request(
request_data=request_data,
route=route,
request_headers=request_headers,
request_query_params=request_query_params,
)
model = _format_model_candidates(candidates)
# If no explicit model was found, try to extract from route
if model is None:
# Parse model from route that follows the pattern /openai/deployments/{model}/*
match = re.match(r"/openai/deployments/([^/]+)", route)

View file

@ -11,7 +11,7 @@ import asyncio
import re
import secrets
from datetime import datetime, timezone
from typing import Any, List, Optional, Tuple, cast
from typing import Any, List, Optional, Tuple, Union, cast
import fastapi
from fastapi import HTTPException, Request, WebSocket, status
@ -63,6 +63,7 @@ from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body,
_safe_get_request_headers,
_safe_get_request_query_params,
populate_request_with_path_params,
)
from litellm.proxy.common_utils.realtime_utils import _realtime_request_body
@ -118,6 +119,29 @@ azure_apim_header = APIKeyHeader(
)
def _get_model_from_request_context(
request_data: dict,
route: str,
request: Optional[Request],
) -> Optional[Union[str, List[str]]]:
return get_model_from_request(
request_data=request_data,
route=route,
request_headers=_safe_get_request_headers(request=request),
request_query_params=_safe_get_request_query_params(request=request),
)
def _get_model_names_for_budget_checks(
model: Optional[Union[str, List[str]]],
) -> List[str]:
if model is None:
return []
if isinstance(model, str):
return [model]
return model
def _get_bearer_token_or_received_api_key(api_key: str) -> str:
if api_key.startswith("Bearer "): # ensure Bearer token passed in
api_key = api_key.replace("Bearer ", "") # extract the token
@ -884,7 +908,11 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
)
# Check if model has zero cost - if so, skip all budget checks
model = get_model_from_request(request_data, route)
model = _get_model_from_request_context(
request_data=request_data,
route=route,
request=request,
)
skip_budget_checks = False
if model is not None and llm_router is not None:
from litellm.proxy.auth.auth_checks import _is_model_cost_zero
@ -1254,6 +1282,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
valid_token=valid_token,
request_data=request_data,
route=route,
request=request,
llm_model_list=llm_model_list,
llm_router=llm_router,
)
@ -1279,7 +1308,11 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
user_obj = None
# Check 2a. Check if model has zero cost - if so, skip all budget checks
model = get_model_from_request(request_data, route)
model = _get_model_from_request_context(
request_data=request_data,
route=route,
request=request,
)
skip_budget_checks = False
if model is not None and llm_router is not None:
from litellm.proxy.auth.auth_checks import _is_model_cost_zero
@ -1403,21 +1436,29 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
# Check 5. Token Model Spend is under Model budget
max_budget_per_model = valid_token.model_max_budget
current_model = request_data.get("model", None)
current_model = _get_model_from_request_context(
request_data=request_data,
route=route,
request=request,
)
current_models = _get_model_names_for_budget_checks(
model=current_model
)
if (
max_budget_per_model is not None
and isinstance(max_budget_per_model, dict)
and len(max_budget_per_model) > 0
and prisma_client is not None
and current_model is not None
and current_models
and valid_token.token is not None
):
## GET THE SPEND FOR THIS MODEL
await model_max_budget_limiter.is_key_within_model_budget(
user_api_key_dict=valid_token,
model=current_model,
)
for model_name in current_models:
await model_max_budget_limiter.is_key_within_model_budget(
user_api_key_dict=valid_token,
model=model_name,
)
# Check 5b. End-user model max budget
end_user_mmb = valid_token.end_user_model_max_budget
@ -1425,14 +1466,15 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
end_user_mmb is not None
and isinstance(end_user_mmb, dict)
and len(end_user_mmb) > 0
and current_model is not None
and current_models
and valid_token.end_user_id is not None
):
await model_max_budget_limiter.is_end_user_within_model_budget(
end_user_id=valid_token.end_user_id,
end_user_model_max_budget=end_user_mmb,
model=current_model,
)
for model_name in current_models:
await model_max_budget_limiter.is_end_user_within_model_budget(
end_user_id=valid_token.end_user_id,
end_user_model_max_budget=end_user_mmb,
model=model_name,
)
# Check 6: Additional Common Checks across jwt + key auth
if valid_token.team_id is not None:
@ -1862,10 +1904,12 @@ async def _run_centralized_common_checks(
user_api_key_auth_obj.project_metadata = project_object.metadata
user_api_key_auth_obj.project_alias = project_object.project_alias
skip_budget_checks = False
model = get_model_from_request(request_data, route)
if model is not None and llm_router is not None:
skip_budget_checks = _is_model_cost_zero(model=model, llm_router=llm_router)
skip_budget_checks = _should_skip_budget_checks(
request_data=request_data,
route=route,
request=request,
llm_router=llm_router,
)
_ = await common_checks(
request=request,
@ -1883,6 +1927,21 @@ async def _run_centralized_common_checks(
project_object=project_object,
)
await _reserve_budget_after_common_checks(
user_api_key_auth_obj=user_api_key_auth_obj,
request_data=request_data,
route=route,
llm_router=llm_router,
team_object=team_object,
user_object=user_object,
end_user_id=end_user_id,
end_user_object=end_user_object,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
skip_budget_checks=skip_budget_checks,
)
async def _noop_none() -> None:
"""Sentinel coroutine for asyncio.gather when a fetch is unnecessary
@ -1890,6 +1949,59 @@ async def _noop_none() -> None:
return None
async def _reserve_budget_after_common_checks(
user_api_key_auth_obj: UserAPIKeyAuth,
request_data: dict,
route: str,
llm_router: Optional[Any],
team_object: Optional[LiteLLM_TeamTableCachedObj],
user_object: Optional[LiteLLM_UserTable],
prisma_client: Optional[PrismaClient],
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
skip_budget_checks: bool,
end_user_id: Optional[str] = None,
end_user_object: Optional[LiteLLM_EndUserTable] = None,
) -> None:
user_api_key_auth_obj.budget_reservation = None
if skip_budget_checks:
return
from litellm.proxy.spend_tracking.budget_reservation import (
reserve_budget_for_request,
)
user_api_key_auth_obj.budget_reservation = await reserve_budget_for_request(
request_body=request_data,
route=route,
llm_router=llm_router,
valid_token=user_api_key_auth_obj,
team_object=team_object,
user_object=user_object,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
end_user_id=end_user_id,
end_user_object=end_user_object,
)
def _should_skip_budget_checks(
request_data: dict,
route: str,
request: Optional[Request],
llm_router: Optional[Any],
) -> bool:
model = _get_model_from_request_context(
request_data=request_data,
route=route,
request=request,
)
if model is not None and llm_router is not None:
return _is_model_cost_zero(model=model, llm_router=llm_router)
return False
@tracer.wrap()
async def user_api_key_auth(
request: Request,
@ -1927,6 +2039,7 @@ async def user_api_key_auth(
request_data=request_data,
custom_litellm_key_header=custom_litellm_key_header,
)
user_api_key_auth_obj.budget_reservation = None
## ENSURE DISABLE ROUTE WORKS ACROSS ALL USER AUTH FLOWS ##
RouteChecks.should_call_route(route=route, valid_token=user_api_key_auth_obj)
@ -2134,6 +2247,7 @@ async def _enforce_key_and_fallback_model_access(
valid_token: UserAPIKeyAuth,
request_data: dict,
route: str,
request: Optional[Request],
llm_model_list: Optional[list],
llm_router: Optional[Any],
) -> None:
@ -2152,7 +2266,11 @@ async def _enforce_key_and_fallback_model_access(
):
pass
else:
model = get_model_from_request(request_data, route)
model = _get_model_from_request_context(
request_data=request_data,
route=route,
request=request,
)
fallback_models = cast(
Optional[List[ALL_FALLBACK_MODEL_VALUES]],
request_data.get("fallbacks", None),
@ -2239,11 +2357,17 @@ async def _run_post_custom_auth_checks(
valid_token=valid_token,
request_data=request_data,
route=route,
request=request,
llm_model_list=llm_model_list,
llm_router=llm_router,
)
current_model = request_data.get("model", None)
current_model = _get_model_from_request_context(
request_data=request_data,
route=route,
request=request,
)
current_models = _get_model_names_for_budget_checks(model=current_model)
# 3. Check key-level model_max_budget
max_budget_per_model = valid_token.model_max_budget
@ -2251,13 +2375,14 @@ async def _run_post_custom_auth_checks(
max_budget_per_model is not None
and isinstance(max_budget_per_model, dict)
and len(max_budget_per_model) > 0
and current_model is not None
and current_models
and valid_token.token is not None
):
await model_max_budget_limiter.is_key_within_model_budget(
user_api_key_dict=valid_token,
model=current_model,
)
for model_name in current_models:
await model_max_budget_limiter.is_key_within_model_budget(
user_api_key_dict=valid_token,
model=model_name,
)
# 4. Check end-user model_max_budget
end_user_mmb = valid_token.end_user_model_max_budget
@ -2265,14 +2390,15 @@ async def _run_post_custom_auth_checks(
end_user_mmb is not None
and isinstance(end_user_mmb, dict)
and len(end_user_mmb) > 0
and current_model is not None
and current_models
and valid_token.end_user_id is not None
):
await model_max_budget_limiter.is_end_user_within_model_budget(
end_user_id=valid_token.end_user_id,
end_user_model_max_budget=end_user_mmb,
model=current_model,
)
for model_name in current_models:
await model_max_budget_limiter.is_end_user_within_model_budget(
end_user_id=valid_token.end_user_id,
end_user_model_max_budget=end_user_mmb,
model=model_name,
)
# team / user / end_user / project context objects are fetched by
# the centralized common_checks gate in user_api_key_auth after

View file

@ -97,6 +97,55 @@ def _serialize_http_exception_detail(
return str(detail), None
def _collect_response_file_search_vector_store_ids(data: Dict[str, Any]) -> set[str]:
vector_store_ids: set[str] = set()
tools = data.get("tools")
if not isinstance(tools, list):
return vector_store_ids
for tool in tools:
if not isinstance(tool, dict) or tool.get("type") != "file_search":
continue
ids = tool.get("vector_store_ids") or []
if not isinstance(ids, list):
raise HTTPException(
status_code=400,
detail={
"error": "file_search.vector_store_ids must be a list of strings"
},
)
for vector_store_id in ids:
if not isinstance(vector_store_id, str) or not vector_store_id:
raise HTTPException(
status_code=400,
detail={
"error": "file_search.vector_store_ids must be a list of strings"
},
)
vector_store_ids.add(vector_store_id)
return vector_store_ids
async def _authorize_response_file_search_vector_stores(
data: Dict[str, Any],
user_api_key_dict: UserAPIKeyAuth,
) -> None:
vector_store_ids = _collect_response_file_search_vector_store_ids(data)
if not vector_store_ids:
return
from litellm.proxy.vector_store_endpoints.utils import (
assert_user_can_access_vector_store_id,
)
for vector_store_id in sorted(vector_store_ids):
await assert_user_can_access_vector_store_id(
vector_store_id=vector_store_id,
user_api_key_dict=user_api_key_dict,
)
async def _parse_event_data_for_error(event_line: Union[str, bytes]) -> Optional[int]:
"""Parses an event line and returns an error code if present, else None."""
event_line = (
@ -791,6 +840,11 @@ class ProxyBaseLLMRequestProcessing:
version=version,
proxy_config=proxy_config,
)
if route_type in {"aresponses", "_aresponses_websocket"}:
await _authorize_response_file_search_vector_stores(
data=self.data,
user_api_key_dict=user_api_key_dict,
)
# Calculate request queue time after add_litellm_data_to_request
# which sets arrival_time in proxy_server_request

View file

@ -14,10 +14,12 @@ memory in long-lived deployments.
import asyncio
from collections import OrderedDict
from datetime import datetime
from typing import TYPE_CHECKING, ClassVar, Optional
from litellm._logging import verbose_proxy_logger
from litellm.constants import SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
if TYPE_CHECKING:
from litellm.caching.dual_cache import DualCache
@ -35,6 +37,10 @@ class SpendCounterReseed:
spend:team_member:{uid}:{tid} -> LiteLLM_TeamMembership.spend
spend:user:{user_id} -> LiteLLM_UserTable.spend
spend:org:{org_id} -> LiteLLM_OrganizationTable.spend
End-user and tag spend counters intentionally do not reseed here. Their
auth paths already load the corresponding objects via get_end_user_object()
and get_tag_objects_batch(); callers pass those values as fallback_spend.
"""
_locks: ClassVar["OrderedDict[str, asyncio.Lock]"] = OrderedDict()
@ -69,9 +75,10 @@ class SpendCounterReseed:
"""
if prisma_client is None:
return None
# Per-window counters share prefixes with primary counters but
# don't correspond to a DB row.
if ":window:" in counter_key:
# Per-window key/team counters share prefixes with primary counters
# but don't correspond to a DB row. Do not reject arbitrary entity IDs
# or tag names that merely contain ":window:".
if SpendCounterReseed._is_key_or_team_window_counter(counter_key):
return None
try:
if counter_key.startswith("spend:key:"):
@ -97,6 +104,10 @@ class SpendCounterReseed:
row = await prisma_client.db.litellm_usertable.find_unique(
where={"user_id": user_id}
)
elif counter_key.startswith("spend:end_user:"):
return None
elif counter_key.startswith("spend:tag:"):
return None
elif counter_key.startswith("spend:org:"):
org_id = counter_key[len("spend:org:") :]
row = await prisma_client.db.litellm_organizationtable.find_unique(
@ -113,11 +124,27 @@ class SpendCounterReseed:
return None
return float(getattr(row, "spend", 0.0) or 0.0)
@staticmethod
def _is_key_or_team_window_counter(counter_key: str) -> bool:
for prefix in ("spend:key:", "spend:team:"):
if not counter_key.startswith(prefix):
continue
_, separator, duration = counter_key.rpartition(":window:")
if not separator or not duration:
return False
try:
duration_in_seconds(duration)
except Exception:
return False
return True
return False
@staticmethod
async def coalesced(
prisma_client: Optional["PrismaClient"],
spend_counter_cache: "DualCache",
counter_key: str,
require_cache_warm: bool = False,
) -> Optional[float]:
"""
Reseed a cold spend counter from the DB and warm the cache,
@ -152,12 +179,156 @@ class SpendCounterReseed:
return None
# Warm even when 0 so subsequent reads hit cache, not DB.
try:
await spend_counter_cache.async_increment_cache(
key=counter_key, value=db_spend, refresh_ttl=True
)
if spend_counter_cache.redis_cache is not None:
current_value = (
await spend_counter_cache.redis_cache.async_increment(
key=counter_key,
value=db_spend,
refresh_ttl=True,
)
)
spend_counter_cache.in_memory_cache.set_cache(
key=counter_key,
value=current_value,
)
else:
await spend_counter_cache.async_increment_cache(
key=counter_key, value=db_spend, refresh_ttl=True
)
except Exception:
verbose_proxy_logger.exception(
"SpendCounterReseed.coalesced: failed to warm counter %s",
counter_key,
)
if require_cache_warm:
raise
return db_spend
@staticmethod
async def window_from_spend_logs(
prisma_client: Optional["PrismaClient"],
entity_type: str,
entity_id: str,
window_start: datetime,
) -> Optional[float]:
if prisma_client is None:
return None
if entity_type == "Key":
group_field = "api_key"
where = {
"api_key": entity_id,
"startTime": {"gte": window_start},
}
elif entity_type == "Team":
group_field = "team_id"
where = {
"team_id": entity_id,
"startTime": {"gte": window_start},
}
else:
return None
try:
response = await prisma_client.db.litellm_spendlogs.group_by(
by=[group_field],
where=where, # type: ignore[arg-type]
sum={"spend": True},
)
except Exception:
verbose_proxy_logger.exception(
"SpendCounterReseed.window_from_spend_logs: failed for %s=%s",
entity_type,
entity_id,
)
return None
if not response:
return 0.0
first_row = response[0]
sum_row = (
first_row.get("_sum")
if isinstance(first_row, dict)
else getattr(first_row, "_sum", None)
)
spend = (
sum_row.get("spend")
if isinstance(sum_row, dict)
else getattr(sum_row, "spend", None)
)
return float(spend or 0.0)
@staticmethod
async def coalesced_window(
prisma_client: Optional["PrismaClient"],
spend_counter_cache: "DualCache",
counter_key: str,
entity_type: str,
entity_id: str,
window_start: datetime,
) -> Optional[float]:
lock = await SpendCounterReseed._get_lock(counter_key)
async with lock:
redis_clean_miss = False
if spend_counter_cache.redis_cache is not None:
try:
val = await spend_counter_cache.redis_cache.async_get_cache(
key=counter_key
)
if val is not None:
return float(val)
redis_clean_miss = True
except Exception:
pass
if not redis_clean_miss:
val = spend_counter_cache.in_memory_cache.get_cache(key=counter_key)
if val is not None:
return float(val)
window_spend = await SpendCounterReseed.window_from_spend_logs(
prisma_client=prisma_client,
entity_type=entity_type,
entity_id=entity_id,
window_start=window_start,
)
if window_spend is None:
return None
try:
if spend_counter_cache.redis_cache is not None:
seeded = await spend_counter_cache.redis_cache.async_set_cache(
key=counter_key,
value=window_spend,
nx=True,
)
if seeded:
current_value = window_spend
else:
current_cached_value = (
await spend_counter_cache.redis_cache.async_get_cache(
key=counter_key
)
)
if current_cached_value is None:
current_value = (
await spend_counter_cache.redis_cache.async_increment(
key=counter_key,
value=window_spend,
)
)
else:
current_value = float(current_cached_value)
spend_counter_cache.in_memory_cache.set_cache(
key=counter_key,
value=current_value,
)
else:
await spend_counter_cache.async_increment_cache(
key=counter_key, value=window_spend
)
except Exception:
verbose_proxy_logger.exception(
"SpendCounterReseed.coalesced_window: failed to warm counter %s",
counter_key,
)
raise
return window_spend

View file

@ -32,10 +32,25 @@ class _PROXY_MaxBudgetLimiter(CustomLogger):
if user_api_key_dict.team_id is not None:
return
# The reservation path admits at the strict-`<` boundary and
# atomically pre-fills the same counter we'd read here. Re-checking
# with `>=` would reject a request the reservation already admitted
# when the reservation fills the counter to exactly max_budget.
# Imported lazily to avoid a circular import via proxy.utils.
from litellm.proxy.spend_tracking.budget_reservation import (
get_reserved_counter_keys,
)
user_counter_key = f"spend:user:{user_id}"
if user_counter_key in get_reserved_counter_keys(
user_api_key_dict.budget_reservation
):
return
from litellm.proxy.proxy_server import get_current_spend
curr_spend = await get_current_spend(
counter_key=f"spend:user:{user_id}",
counter_key=user_counter_key,
fallback_spend=user_api_key_dict.user_spend or 0.0,
)

View file

@ -30,16 +30,35 @@ class _ProxyDBLogger(CustomLogger):
kwargs, response_obj, start_time, end_time
)
async def async_post_call_failure_hook(
self,
request_data: dict,
original_exception: Exception,
user_api_key_dict: UserAPIKeyAuth,
traceback_str: Optional[str] = None,
):
request_route = user_api_key_dict.request_route
if _ProxyDBLogger._should_track_errors_in_db() is False:
return
async def async_post_call_failure_hook(
self,
request_data: dict,
original_exception: Exception,
user_api_key_dict: UserAPIKeyAuth,
traceback_str: Optional[str] = None,
):
try:
await _release_budget_reservation(
budget_reservation=user_api_key_dict.budget_reservation
)
except Exception:
verbose_proxy_logger.exception(
"Failed to release budget reservation during failure handling"
)
try:
await _invalidate_budget_reservation_counters(
budget_reservation=user_api_key_dict.budget_reservation
)
if user_api_key_dict.budget_reservation is not None:
user_api_key_dict.budget_reservation["finalized"] = True
except Exception:
verbose_proxy_logger.exception(
"Failed to invalidate budget reservation counters after failure release failed"
)
request_route = user_api_key_dict.request_route
if _ProxyDBLogger._should_track_errors_in_db() is False:
return
elif request_route is not None and not (
RouteChecks.is_llm_api_route(route=request_route)
or RouteChecks.is_info_route(route=request_route)
@ -155,66 +174,64 @@ class _ProxyDBLogger(CustomLogger):
f"kwargs stream: {kwargs.get('stream', None)} + complete streaming response: {kwargs.get('complete_streaming_response', None)}"
)
parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs=kwargs)
litellm_params = kwargs.get("litellm_params", {}) or {}
end_user_id = get_end_user_id_for_cost_tracking(litellm_params)
metadata = get_litellm_metadata_from_kwargs(kwargs=kwargs)
user_id = cast(Optional[str], metadata.get("user_api_key_user_id", None))
team_id = cast(Optional[str], metadata.get("user_api_key_team_id", None))
org_id = cast(Optional[str], metadata.get("user_api_key_org_id", None))
litellm_params = kwargs.get("litellm_params", {}) or {}
end_user_id = get_end_user_id_for_cost_tracking(litellm_params)
metadata = get_litellm_metadata_from_kwargs(kwargs=kwargs)
budget_reservation = _get_budget_reservation_from_metadata(
metadata=metadata
)
user_id = cast(Optional[str], metadata.get("user_api_key_user_id", None))
team_id = cast(Optional[str], metadata.get("user_api_key_team_id", None))
org_id = cast(Optional[str], metadata.get("user_api_key_org_id", None))
key_alias = cast(Optional[str], metadata.get("user_api_key_alias", None))
end_user_max_budget = metadata.get("user_api_end_user_max_budget", None)
sl_object: Optional[StandardLoggingPayload] = kwargs.get(
"standard_logging_object", None
)
response_cost = (
sl_object.get("response_cost", None)
if sl_object is not None
else kwargs.get("response_cost", None)
)
tags: Optional[List[str]] = (
sl_object.get("request_tags", None) if sl_object is not None else None
)
if response_cost is not None:
user_api_key = metadata.get("user_api_key", None)
response_cost = (
sl_object.get("response_cost", None)
if sl_object is not None
else kwargs.get("response_cost", None)
)
tags = _get_request_tags_for_cost_tracking(
sl_object=sl_object,
metadata=metadata,
)
if response_cost is not None:
user_api_key = metadata.get("user_api_key", None)
if kwargs.get("cache_hit", False) is True:
response_cost = 0.0
verbose_proxy_logger.debug(
f"Cache Hit: response_cost {response_cost}, for user_id {user_id}"
)
verbose_proxy_logger.debug(
f"user_api_key {user_api_key}, user_id {user_id}, team_id {team_id}, end_user_id {end_user_id}"
)
if _should_track_cost_callback(
user_api_key=user_api_key,
verbose_proxy_logger.debug(
f"user_api_key {user_api_key}, user_id {user_id}, team_id {team_id}, end_user_id {end_user_id}"
)
if _should_track_cost_callback(
user_api_key=user_api_key,
user_id=user_id,
team_id=team_id,
end_user_id=end_user_id,
):
## UPDATE DATABASE
await proxy_logging_obj.db_spend_update_writer.update_database(
token=user_api_key,
response_cost=response_cost,
user_id=user_id,
end_user_id=end_user_id,
team_id=team_id,
kwargs=kwargs,
completion_response=completion_response,
start_time=start_time,
end_time=end_time,
org_id=org_id,
)
# Atomically update spend counters (in-memory + Redis)
# for cross-pod budget enforcement.
await increment_spend_counters(
token=user_api_key,
team_id=team_id,
user_id=user_id,
response_cost=response_cost,
org_id=org_id,
)
end_user_id=end_user_id,
):
## UPDATE DATABASE
await _update_database_and_spend_counters(
proxy_logging_obj=proxy_logging_obj,
increment_spend_counters=increment_spend_counters,
user_api_key=user_api_key,
user_id=user_id,
end_user_id=end_user_id,
team_id=team_id,
org_id=org_id,
kwargs=kwargs,
completion_response=completion_response,
start_time=start_time,
end_time=end_time,
response_cost=response_cost,
budget_reservation=budget_reservation,
request_tags=tags,
)
# update cache (fire-and-forget for backward compat:
# cached object fields, soft budget alerts, etc.)
@ -234,10 +251,15 @@ class _ProxyDBLogger(CustomLogger):
token=user_api_key,
key_alias=key_alias,
end_user_id=end_user_id,
response_cost=response_cost,
max_budget=end_user_max_budget,
)
response_cost=response_cost,
max_budget=end_user_max_budget,
)
elif budget_reservation is not None:
await _release_budget_reservation(
budget_reservation=budget_reservation
)
else:
await _release_budget_reservation(budget_reservation=budget_reservation)
# Non-model call types (health checks, afile_delete) have no model or standard_logging_object.
# Use .get() for "stream" to avoid KeyError on health checks.
if sl_object is None and not kwargs.get("model"):
@ -366,7 +388,7 @@ class _ProxyDBLogger(CustomLogger):
return
def _should_track_cost_callback(
def _should_track_cost_callback(
user_api_key: Optional[str],
user_id: Optional[str],
team_id: Optional[str],
@ -387,4 +409,135 @@ def _should_track_cost_callback(
or end_user_id is not None
):
return True
return False
return False
def _get_budget_reservation_from_metadata(metadata: dict) -> Optional[dict]:
metadata_budget_reservation = metadata.get("user_api_key_budget_reservation")
if isinstance(metadata_budget_reservation, dict):
return metadata_budget_reservation
user_api_key_auth_obj = metadata.get("user_api_key_auth")
if user_api_key_auth_obj is None:
return None
if isinstance(user_api_key_auth_obj, dict):
budget_reservation = user_api_key_auth_obj.get("budget_reservation")
return budget_reservation if isinstance(budget_reservation, dict) else None
return getattr(user_api_key_auth_obj, "budget_reservation", None)
def _get_request_tags_for_cost_tracking(
sl_object: Optional[StandardLoggingPayload],
metadata: dict,
) -> Optional[List[str]]:
if sl_object is not None:
request_tags = sl_object.get("request_tags", None)
if isinstance(request_tags, list):
return request_tags
metadata_tags = metadata.get("tags", None)
if isinstance(metadata_tags, list):
return metadata_tags
return None
async def _update_database_and_spend_counters(
proxy_logging_obj: Any,
increment_spend_counters: Any,
user_api_key: Optional[str],
user_id: Optional[str],
end_user_id: Optional[str],
team_id: Optional[str],
org_id: Optional[str],
kwargs: dict,
completion_response: Optional[Union[litellm.ModelResponse, Any]],
start_time: Any,
end_time: Any,
response_cost: float,
budget_reservation: Optional[dict],
request_tags: Optional[List[str]] = None,
) -> None:
try:
await proxy_logging_obj.db_spend_update_writer.update_database(
token=user_api_key,
response_cost=response_cost,
user_id=user_id,
end_user_id=end_user_id,
team_id=team_id,
kwargs=kwargs,
completion_response=completion_response,
start_time=start_time,
end_time=end_time,
org_id=org_id,
)
except Exception:
if budget_reservation is not None:
try:
await _release_budget_reservation(budget_reservation=budget_reservation)
except Exception:
verbose_proxy_logger.exception(
"Failed to release budget reservation after database update failed"
)
try:
await _invalidate_budget_reservation_counters(
budget_reservation=budget_reservation
)
except Exception:
verbose_proxy_logger.exception(
"Failed to invalidate budget reservation counters after release failed"
)
raise
try:
await increment_spend_counters(
token=user_api_key,
team_id=team_id,
user_id=user_id,
response_cost=response_cost,
org_id=org_id,
budget_reservation=budget_reservation,
end_user_id=end_user_id,
tags=request_tags,
)
except Exception:
if budget_reservation is not None:
try:
await _invalidate_budget_reservation_counters(
budget_reservation=budget_reservation
)
except Exception:
verbose_proxy_logger.exception(
"Failed to invalidate budget reservation counters after spend counter update failed"
)
finally:
budget_reservation["finalized"] = True
raise
async def _release_budget_reservation(budget_reservation: Optional[dict]) -> None:
if budget_reservation is None:
return
from litellm.proxy.spend_tracking.budget_reservation import (
release_budget_reservation,
)
await release_budget_reservation(
budget_reservation=budget_reservation,
)
async def _invalidate_budget_reservation_counters(
budget_reservation: Optional[dict],
) -> None:
if budget_reservation is None:
return
from litellm.proxy.spend_tracking.budget_reservation import (
invalidate_budget_reservation_counters,
)
await invalidate_budget_reservation_counters(
budget_reservation=budget_reservation,
)

View file

@ -893,6 +893,10 @@ class LiteLLMProxyRequestSetup:
data[_metadata_variable_name]["user_api_end_user_max_budget"] = getattr(
user_api_key_dict, "end_user_max_budget", None
)
if user_api_key_dict.budget_reservation is not None:
data[_metadata_variable_name][
"user_api_key_budget_reservation"
] = user_api_key_dict.budget_reservation
# Add the full UserAPIKeyAuth object for MCP server access control
data[_metadata_variable_name]["user_api_key_auth"] = user_api_key_dict
return data

View file

@ -47,6 +47,8 @@ from litellm.types.passthrough_endpoints.pass_through_endpoints import (
)
from litellm.proxy.utils import is_known_model
from litellm.proxy.vector_store_endpoints.utils import (
assert_user_can_access_vector_store,
get_litellm_managed_vector_store,
is_allowed_to_call_vector_store_endpoint,
)
from litellm.secret_managers.main import get_secret_str
@ -533,6 +535,10 @@ async def milvus_proxy_route(
)
if vector_store is None:
raise Exception(f"Vector store not found for {vector_store_name}")
await assert_user_can_access_vector_store(
vector_store=vector_store,
user_api_key_dict=user_api_key_dict,
)
litellm_params = vector_store.get("litellm_params") or {}
auth_credentials = provider_config.get_auth_credentials(
litellm_params=litellm_params
@ -1438,6 +1444,10 @@ async def azure_proxy_route(
)
if vector_store is None:
raise Exception(f"Vector store not found for {vector_store_name}")
await assert_user_can_access_vector_store(
vector_store=vector_store,
user_api_key_dict=user_api_key_dict,
)
litellm_params = vector_store.get("litellm_params") or {}
auth_credentials = provider_config.get_auth_credentials(
litellm_params=litellm_params
@ -1777,6 +1787,11 @@ async def _base_vertex_proxy_route(
request=request,
api_key=api_key_to_use,
)
if router_credentials is not None:
await assert_user_can_access_vector_store(
vector_store=router_credentials,
user_api_key_dict=user_api_key_dict,
)
vertex_project: Optional[str] = get_vertex_project_id_from_url(endpoint)
vertex_location: Optional[str] = get_vertex_location_from_url(endpoint)
@ -1913,11 +1928,11 @@ async def vertex_discovery_proxy_route(
"Extracted vector store ID from endpoint: %s", vector_store_id
)
# Retrieve vector store credentials from the registry
vector_store_credentials = (
passthrough_endpoint_router.get_vector_store_credentials(
vector_store_id=vector_store_id
)
# Retrieve LiteLLM-managed vector store credentials if the datastore id
# is registered with LiteLLM. Unknown datastore ids keep the existing
# direct Vertex pass-through behavior.
vector_store_credentials = await get_litellm_managed_vector_store(
vector_store_id=vector_store_id
)
if vector_store_credentials:
@ -1925,7 +1940,7 @@ async def vertex_discovery_proxy_route(
"Found vector store credentials for ID: %s", vector_store_id
)
else:
verbose_proxy_logger.warning(
verbose_proxy_logger.debug(
"Vector store ID %s found in endpoint but no credentials found in registry",
vector_store_id,
)

View file

@ -2324,14 +2324,10 @@ async def _register_pass_through_endpoint(
dependencies = None
if auth is not None and str(auth).lower() == "true":
# Authentication on a pass-through endpoint used to be enterprise-
# only — which left the OSS tier with no safe configuration: the
# default was ``auth=False`` (unauthenticated forwarder) and the
# safe ``auth=True`` raised at startup unless the operator had a
# license. The default is now ``True`` (safe-by-default), and
# turning it on no longer requires a license: an unauthenticated
# forwarder is a deployment choice the operator should be allowed
# to make explicitly, but the safe option must always be free.
# Authentication on a pass-through endpoint used to be enterprise-only.
# That left OSS with no safe configuration: auth=True raised at startup
# unless the operator had a license. The safe option must always be free,
# and unauthenticated forwarding should require explicit opt-in.
dependencies = [Depends(user_api_key_auth)]
if path not in LiteLLMRoutes.openai_routes.value:
LiteLLMRoutes.openai_routes.value.append(path)

View file

@ -6,6 +6,7 @@ import inspect
import io
import os
import random
import re
import secrets
import shutil
import subprocess
@ -955,6 +956,85 @@ async def proxy_startup_event(app: FastAPI): # noqa: PLR0915
await proxy_shutdown_event() # type: ignore[reportGeneralTypeIssues]
def _generate_stable_operation_id(route: Any) -> str:
operation_id = re.sub(r"\W", "_", f"{route.name}{route.path_format}")
route_methods = sorted(route.methods or [])
if len(route_methods) == 1:
operation_id = f"{operation_id}_{route_methods[0].lower()}"
return operation_id
_OPENAPI_HTTP_METHODS = {
"delete",
"get",
"head",
"options",
"patch",
"post",
"put",
"trace",
}
def _strip_operation_id_method_suffix(operation_id: str) -> str:
base, separator, suffix = operation_id.rpartition("_")
if separator and suffix in _OPENAPI_HTTP_METHODS:
return base
return operation_id
def ensure_unique_openapi_operation_ids(
openapi_schema: Dict[str, Any],
reserved_operation_ids: Optional[Set[str]] = None,
) -> Dict[str, Any]:
operation_entries = []
operation_id_counts: Dict[str, int] = {}
for path_item in openapi_schema.get("paths", {}).values():
if not isinstance(path_item, dict):
continue
for method, operation in path_item.items():
if method not in _OPENAPI_HTTP_METHODS or not isinstance(operation, dict):
continue
operation_id = operation.get("operationId")
if not isinstance(operation_id, str):
continue
operation_entries.append((method, operation, operation_id))
operation_id_counts[operation_id] = (
operation_id_counts.get(operation_id, 0) + 1
)
used_operation_ids = set(reserved_operation_ids or set())
seen_operation_ids: Set[str] = set()
for method, operation, operation_id in operation_entries:
should_rewrite = (
operation_id_counts[operation_id] > 1
or operation_id in used_operation_ids
or operation_id in seen_operation_ids
)
if not should_rewrite:
seen_operation_ids.add(operation_id)
used_operation_ids.add(operation_id)
continue
base_operation_id = _strip_operation_id_method_suffix(operation_id)
new_operation_id = f"{base_operation_id}_{method}"
suffix = 2
while (
new_operation_id in used_operation_ids
or new_operation_id in seen_operation_ids
):
new_operation_id = f"{base_operation_id}_{method}_{suffix}"
suffix += 1
operation["operationId"] = new_operation_id
seen_operation_ids.add(new_operation_id)
used_operation_ids.add(new_operation_id)
if reserved_operation_ids is not None:
reserved_operation_ids.update(used_operation_ids)
return openapi_schema
app = FastAPI(
docs_url=_get_docs_url(),
redoc_url=_get_redoc_url(),
@ -964,6 +1044,7 @@ app = FastAPI(
version=version,
root_path=server_root_path,
lifespan=proxy_startup_event, # type: ignore[reportGeneralTypeIssues]
generate_unique_id_function=_generate_stable_operation_id,
)
vertex_live_passthrough_vertex_base = VertexBase()
@ -1043,6 +1124,7 @@ def get_openapi_schema():
from litellm.proxy._lazy_features import inject_lazy_stubs
openapi_schema = inject_lazy_stubs(openapi_schema)
openapi_schema = ensure_unique_openapi_operation_ids(openapi_schema)
# Fix Swagger UI execute path error when server_root_path is set
if server_root_path:
@ -1074,6 +1156,7 @@ def custom_openapi():
from litellm.proxy._lazy_features import inject_lazy_stubs
openapi_schema = inject_lazy_stubs(openapi_schema)
openapi_schema = ensure_unique_openapi_operation_ids(openapi_schema)
# Fix Swagger UI execute path error when server_root_path is set
if server_root_path:
@ -1845,6 +1928,9 @@ async def increment_spend_counters(
user_id: Optional[str],
response_cost: Optional[float],
org_id: Optional[str] = None,
budget_reservation: Optional[dict] = None,
end_user_id: Optional[str] = None,
tags: Optional[List[str]] = None,
):
"""
Atomically increment spend counters for budget enforcement.
@ -1856,7 +1942,14 @@ async def increment_spend_counters(
Awaited (not create_task) in the cost callback, so the counter is
updated before the next request's auth check runs.
"""
reserved_counter_keys = await _reconcile_budget_reservation_for_counter_update(
budget_reservation=budget_reservation,
response_cost=response_cost,
)
if response_cost is None or response_cost == 0:
if budget_reservation is not None:
budget_reservation["finalized"] = True
return
if token is not None:
@ -1871,11 +1964,13 @@ async def increment_spend_counters(
if isinstance(token, str) and token.startswith("sk-")
else token
)
await _init_and_increment_spend_counter(
counter_key=f"spend:key:{hashed_token}",
source_cache_key=hashed_token,
increment=response_cost,
)
key_counter_key = f"spend:key:{hashed_token}"
if key_counter_key not in reserved_counter_keys:
await _init_and_increment_spend_counter(
counter_key=key_counter_key,
source_cache_key=hashed_token,
increment=response_cost,
)
# Increment per-window budget counters for multi-budget keys
key_obj = await user_api_key_cache.async_get_cache(key=hashed_token)
@ -1892,17 +1987,28 @@ async def increment_spend_counters(
if isinstance(window, dict)
else window.budget_duration
)
await spend_counter_cache.async_increment_cache(
key=f"spend:key:{hashed_token}:window:{duration}",
value=response_cost,
)
key_window_counter = f"spend:key:{hashed_token}:window:{duration}"
if key_window_counter not in reserved_counter_keys:
from litellm.proxy.spend_tracking.budget_reservation import (
get_budget_window_start,
)
await _init_and_increment_window_spend_counter(
counter_key=key_window_counter,
entity_type="Key",
entity_id=hashed_token,
window_start=get_budget_window_start(window),
increment=response_cost,
)
if team_id is not None:
await _init_and_increment_spend_counter(
counter_key=f"spend:team:{team_id}",
source_cache_key=f"team_id:{team_id}",
increment=response_cost,
)
team_counter_key = f"spend:team:{team_id}"
if team_counter_key not in reserved_counter_keys:
await _init_and_increment_spend_counter(
counter_key=team_counter_key,
source_cache_key=f"team_id:{team_id}",
increment=response_cost,
)
# Increment per-window budget counters for multi-budget teams
team_obj = await user_api_key_cache.async_get_cache(key=f"team_id:{team_id}")
@ -1919,36 +2025,157 @@ async def increment_spend_counters(
if isinstance(window, dict)
else window.budget_duration
)
await spend_counter_cache.async_increment_cache(
key=f"spend:team:{team_id}:window:{duration}",
value=response_cost,
)
team_window_counter = f"spend:team:{team_id}:window:{duration}"
if team_window_counter not in reserved_counter_keys:
from litellm.proxy.spend_tracking.budget_reservation import (
get_budget_window_start,
)
await _init_and_increment_window_spend_counter(
counter_key=team_window_counter,
entity_type="Team",
entity_id=team_id,
window_start=get_budget_window_start(window),
increment=response_cost,
)
if user_id is not None and team_id is not None:
await _init_and_increment_spend_counter(
counter_key=f"spend:team_member:{user_id}:{team_id}",
source_cache_key=f"team_membership:{user_id}:{team_id}",
increment=response_cost,
)
team_member_counter_key = f"spend:team_member:{user_id}:{team_id}"
if team_member_counter_key not in reserved_counter_keys:
await _init_and_increment_spend_counter(
counter_key=team_member_counter_key,
source_cache_key=f"team_membership:{user_id}:{team_id}",
increment=response_cost,
)
if user_id is not None:
await _init_and_increment_spend_counter(
counter_key=f"spend:user:{user_id}",
source_cache_key=user_id,
user_counter_key = f"spend:user:{user_id}"
if user_counter_key not in reserved_counter_keys:
await _init_and_increment_spend_counter(
counter_key=user_counter_key,
source_cache_key=user_id,
increment=response_cost,
)
await _increment_end_user_and_tag_spend_counters(
end_user_id=end_user_id,
tags=tags,
response_cost=response_cost,
reserved_counter_keys=reserved_counter_keys,
)
await _increment_org_spend_counter(
org_id=org_id,
response_cost=response_cost,
reserved_counter_keys=reserved_counter_keys,
)
if budget_reservation is not None:
budget_reservation["finalized"] = True
async def _reconcile_budget_reservation_for_counter_update(
budget_reservation: Optional[dict],
response_cost: Optional[float],
) -> Set[str]:
if budget_reservation is None:
return set()
from litellm.proxy.spend_tracking.budget_reservation import (
get_reserved_counter_keys,
invalidate_budget_reservation_counters,
reconcile_budget_reservation,
)
reserved_counter_keys = get_reserved_counter_keys(
budget_reservation=budget_reservation
)
try:
await reconcile_budget_reservation(
budget_reservation=budget_reservation,
actual_cost=response_cost or 0.0,
finalize=False,
)
except Exception:
verbose_proxy_logger.warning(
"Failed to reconcile budget reservation after persisted spend; invalidating reserved counters and continuing",
exc_info=True,
)
try:
await invalidate_budget_reservation_counters(
budget_reservation=budget_reservation
)
except Exception:
verbose_proxy_logger.exception(
"Failed to invalidate reserved counters after reservation reconciliation failed"
)
return reserved_counter_keys
async def _increment_end_user_and_tag_spend_counters(
end_user_id: Optional[str],
tags: Optional[List[str]],
response_cost: float,
reserved_counter_keys: Set[str],
) -> None:
if end_user_id is not None:
await _init_and_increment_unreserved_spend_counter(
counter_key=f"spend:end_user:{end_user_id}",
source_cache_key=f"end_user_id:{end_user_id}",
increment=response_cost,
reserved_counter_keys=reserved_counter_keys,
)
if org_id is not None:
await _init_and_increment_spend_counter(
counter_key=f"spend:org:{org_id}",
source_cache_key=f"org_id:{org_id}",
if tags is None:
return
seen_tags: Set[str] = set()
for tag_name in tags:
if not tag_name or not isinstance(tag_name, str) or tag_name in seen_tags:
continue
seen_tags.add(tag_name)
await _init_and_increment_unreserved_spend_counter(
counter_key=f"spend:tag:{tag_name}",
source_cache_key=f"tag:{tag_name}",
increment=response_cost,
reserved_counter_keys=reserved_counter_keys,
)
async def _increment_org_spend_counter(
org_id: Optional[str],
response_cost: float,
reserved_counter_keys: Set[str],
) -> None:
if org_id is None:
return
await _init_and_increment_unreserved_spend_counter(
counter_key=f"spend:org:{org_id}",
source_cache_key=[f"org_id:{org_id}:with_budget", f"org_id:{org_id}"],
increment=response_cost,
reserved_counter_keys=reserved_counter_keys,
)
async def _init_and_increment_unreserved_spend_counter(
counter_key: str,
source_cache_key: Union[str, List[str]],
increment: float,
reserved_counter_keys: Set[str],
) -> None:
if counter_key in reserved_counter_keys:
return
await _init_and_increment_spend_counter(
counter_key=counter_key,
source_cache_key=source_cache_key,
increment=increment,
)
async def _init_and_increment_spend_counter(
counter_key: str,
source_cache_key: str,
source_cache_key: Union[str, List[str]],
increment: float,
):
"""
@ -1967,31 +2194,163 @@ async def _init_and_increment_spend_counter(
under-counting (would allow overspend).
4. Increment atomically (both in-memory + Redis)
"""
current = await spend_counter_cache.async_get_cache(key=counter_key)
if current is None:
await _ensure_spend_counter_initialized(
counter_key=counter_key,
source_cache_key=source_cache_key,
)
await _increment_spend_counter_cache(counter_key=counter_key, increment=increment)
async def _init_and_increment_window_spend_counter(
counter_key: str,
entity_type: str,
entity_id: str,
window_start: Optional[datetime],
increment: float,
):
if window_start is None:
verbose_proxy_logger.warning(
"Skipping spend counter increment for invalid budget window %s",
counter_key,
)
return
initialized = await _ensure_window_spend_counter_initialized(
counter_key=counter_key,
entity_type=entity_type,
entity_id=entity_id,
window_start=window_start,
)
if initialized is False:
return
await _increment_spend_counter_cache(counter_key=counter_key, increment=increment)
async def _ensure_spend_counter_initialized(
counter_key: str,
source_cache_key: Union[str, List[str]],
):
is_warm = await _is_spend_counter_cache_warm(counter_key=counter_key)
if is_warm is False:
# Shares the per-counter lock with get_current_spend.
db_spend = await SpendCounterReseed.coalesced(
prisma_client=prisma_client,
spend_counter_cache=spend_counter_cache,
counter_key=counter_key,
require_cache_warm=True,
)
if db_spend is None:
# DB unavailable - fall back to in-process cache (may be stale).
source = await user_api_key_cache.async_get_cache(key=source_cache_key)
base_spend: float = 0.0
if source is not None:
if isinstance(source, dict):
base_spend = source.get("spend", 0.0) or 0.0
else:
base_spend = getattr(source, "spend", 0.0) or 0.0
base_spend = await _get_source_cache_base_spend(
source_cache_key=source_cache_key
)
if base_spend > 0:
await spend_counter_cache.async_increment_cache(
key=counter_key, value=base_spend, refresh_ttl=True
await _increment_spend_counter_cache(
counter_key=counter_key, increment=base_spend
)
await spend_counter_cache.async_increment_cache(
key=counter_key, value=increment, refresh_ttl=True
async def _get_source_cache_base_spend(
source_cache_key: Union[str, List[str]],
) -> float:
source_cache_keys = (
[source_cache_key] if isinstance(source_cache_key, str) else source_cache_key
)
for cache_key in source_cache_keys:
source = await user_api_key_cache.async_get_cache(key=cache_key)
if source is None:
continue
if isinstance(source, dict):
return float(source.get("spend", 0.0) or 0.0)
return float(getattr(source, "spend", 0.0) or 0.0)
return 0.0
async def _ensure_window_spend_counter_initialized(
counter_key: str,
entity_type: str,
entity_id: str,
window_start: datetime,
) -> bool:
is_warm = await _is_spend_counter_cache_warm(counter_key=counter_key)
if is_warm is True:
return True
window_spend = await SpendCounterReseed.coalesced_window(
prisma_client=prisma_client,
spend_counter_cache=spend_counter_cache,
counter_key=counter_key,
entity_type=entity_type,
entity_id=entity_id,
window_start=window_start,
)
if window_spend is None:
verbose_proxy_logger.warning(
"Skipping cold spend counter seed for %s because window spend could not be loaded",
counter_key,
)
return False
return True
async def _is_spend_counter_cache_warm(counter_key: str) -> bool:
if spend_counter_cache.redis_cache is not None:
try:
current_value = await spend_counter_cache.redis_cache.async_get_cache(
key=counter_key,
)
if current_value is None:
return False
spend_counter_cache.in_memory_cache.set_cache(
key=counter_key,
value=current_value,
)
return True
except Exception as e:
verbose_proxy_logger.debug(
"Unable to read Redis spend counter %s before initialization, falling back to in-memory: %s",
counter_key,
e,
)
return spend_counter_cache.in_memory_cache.get_cache(key=counter_key) is not None
async def _increment_spend_counter_cache(counter_key: str, increment: float):
if spend_counter_cache.redis_cache is not None:
try:
current_value = await spend_counter_cache.redis_cache.async_increment(
key=counter_key,
value=increment,
refresh_ttl=True,
)
except Exception:
await _invalidate_spend_counter(counter_key=counter_key)
raise
spend_counter_cache.in_memory_cache.set_cache(
key=counter_key,
value=current_value,
)
return current_value
return await spend_counter_cache.async_increment_cache(
key=counter_key,
value=increment,
refresh_ttl=True,
)
async def _invalidate_spend_counter(counter_key: str):
spend_counter_cache.in_memory_cache.delete_cache(key=counter_key)
if spend_counter_cache.redis_cache is not None:
try:
await spend_counter_cache.redis_cache.async_delete_cache(key=counter_key)
except Exception:
verbose_proxy_logger.debug(
"Unable to delete stale spend counter %s after increment failure",
counter_key,
exc_info=True,
)
async def update_cache( # noqa: PLR0915

View file

@ -15,6 +15,7 @@ from fastapi.responses import ORJSONResponse
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
from litellm.proxy.common_utils.http_parsing_utils import (
@ -22,10 +23,88 @@ from litellm.proxy.common_utils.http_parsing_utils import (
_safe_get_request_headers,
get_form_data,
)
from litellm.proxy.vector_store_endpoints.utils import (
assert_user_can_access_vector_store_id,
)
router = APIRouter()
def _raise_vector_store_scan_depth_exceeded() -> None:
raise HTTPException(
status_code=400,
detail={
"error": f"Max depth of {DEFAULT_MAX_RECURSE_DEPTH} exceeded while scanning vector_store_id values"
},
)
def _append_payload_to_scan_stack(
payload_stack: list[tuple[Any, int]],
value: Any,
next_depth: int,
) -> None:
if isinstance(value, dict):
if next_depth > DEFAULT_MAX_RECURSE_DEPTH:
_raise_vector_store_scan_depth_exceeded()
payload_stack.append((value, next_depth))
elif isinstance(value, list):
if next_depth > DEFAULT_MAX_RECURSE_DEPTH:
if any(isinstance(item, (dict, list)) for item in value):
_raise_vector_store_scan_depth_exceeded()
return
payload_stack.append((value, next_depth))
def _collect_vector_store_ids_from_payload(payload: Any) -> set[str]:
vector_store_ids: set[str] = set()
payload_stack = [(payload, 0)]
while payload_stack:
current_payload, depth = payload_stack.pop()
if depth > DEFAULT_MAX_RECURSE_DEPTH:
_raise_vector_store_scan_depth_exceeded()
if isinstance(current_payload, dict):
for key, value in current_payload.items():
if key == "vector_store_id":
if not isinstance(value, str) or not value:
raise HTTPException(
status_code=400,
detail={
"error": "vector_store_id must be a non-empty string"
},
)
vector_store_ids.add(value)
continue
if isinstance(value, (dict, list)):
_append_payload_to_scan_stack(
payload_stack=payload_stack,
value=value,
next_depth=depth + 1,
)
elif isinstance(current_payload, list):
for item in current_payload:
_append_payload_to_scan_stack(
payload_stack=payload_stack,
value=item,
next_depth=depth + 1,
)
return vector_store_ids
async def _authorize_nested_vector_store_ids(
payload: Any,
user_api_key_dict: UserAPIKeyAuth,
) -> None:
for vector_store_id in sorted(_collect_vector_store_ids_from_payload(payload)):
await assert_user_can_access_vector_store_id(
vector_store_id=vector_store_id,
user_api_key_dict=user_api_key_dict,
)
def _build_file_metadata_entry(
response: Any,
file_data: Optional[Tuple[str, bytes, str]] = None,
@ -385,6 +464,11 @@ async def rag_ingest(
},
)
await _authorize_nested_vector_store_ids(
payload=ingest_options,
user_api_key_dict=user_api_key_dict,
)
# Add litellm data
request_data: Dict[str, Any] = {}
request_data = await add_litellm_data_to_request(
@ -537,11 +621,20 @@ async def rag_query(
status_code=400,
detail={"error": "retrieval_config is required"},
)
if not isinstance(retrieval_config, dict):
raise HTTPException(
status_code=400,
detail={"error": "retrieval_config must be an object"},
)
if "vector_store_id" not in retrieval_config:
raise HTTPException(
status_code=400,
detail={"error": "retrieval_config must contain 'vector_store_id'"},
)
await _authorize_nested_vector_store_ids(
payload=retrieval_config,
user_api_key_dict=user_api_key_dict,
)
# Add litellm data
request_data: Dict[str, Any] = {}

File diff suppressed because it is too large Load diff

View file

@ -1,8 +1,6 @@
from typing import Any, Dict, Optional
from fastapi import APIRouter, Depends, HTTPException, Request, Response
import litellm
from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import (
LiteLLM_ManagedVectorStore,
)
@ -10,7 +8,10 @@ from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.utils import jsonify_object
from litellm.proxy.vector_store_endpoints.utils import can_user_access_vector_store
from litellm.proxy.vector_store_endpoints.utils import (
assert_user_can_access_vector_store,
get_litellm_managed_vector_store,
)
from litellm.types.vector_stores import IndexCreateRequest
router = APIRouter()
@ -19,24 +20,6 @@ router = APIRouter()
########################################################
async def _check_vector_store_access(
vector_store: LiteLLM_ManagedVectorStore,
user_api_key_dict: UserAPIKeyAuth,
) -> bool:
"""
Check if the user has access to the vector store.
Delegates to :func:`can_user_access_vector_store`, which honors:
- PROXY_ADMIN bypass
- legacy vector stores with no team_id
- key-level and team-level ``object_permission.vector_stores`` allowlists
- team_id match between key and store
"""
return await can_user_access_vector_store(
vector_store=vector_store, user_api_key_dict=user_api_key_dict
)
async def _update_request_data_with_litellm_managed_vector_store_registry(
data: Dict,
vector_store_id: str,
@ -53,35 +36,27 @@ async def _update_request_data_with_litellm_managed_vector_store_registry(
Raises:
HTTPException: If user doesn't have access to the vector store
"""
if litellm.vector_store_registry is not None:
vector_store_to_run: Optional[LiteLLM_ManagedVectorStore] = (
litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry(
vector_store_id=vector_store_id
vector_store_to_run: Optional[LiteLLM_ManagedVectorStore] = (
await get_litellm_managed_vector_store(vector_store_id=vector_store_id)
)
if vector_store_to_run is not None:
if user_api_key_dict is not None:
await assert_user_can_access_vector_store(
vector_store=vector_store_to_run,
user_api_key_dict=user_api_key_dict,
)
)
if vector_store_to_run is not None:
if user_api_key_dict is not None:
if not await _check_vector_store_access(
vector_store_to_run, user_api_key_dict
):
raise HTTPException(
status_code=403,
detail="Access denied: You do not have permission to access this vector store",
)
if "custom_llm_provider" in vector_store_to_run:
data["custom_llm_provider"] = vector_store_to_run.get(
"custom_llm_provider"
)
if "custom_llm_provider" in vector_store_to_run:
data["custom_llm_provider"] = vector_store_to_run.get("custom_llm_provider")
if "litellm_credential_name" in vector_store_to_run:
data["litellm_credential_name"] = vector_store_to_run.get(
"litellm_credential_name"
)
if "litellm_credential_name" in vector_store_to_run:
data["litellm_credential_name"] = vector_store_to_run.get(
"litellm_credential_name"
)
if "litellm_params" in vector_store_to_run:
litellm_params = vector_store_to_run.get("litellm_params", {}) or {}
data.update(litellm_params)
if "litellm_params" in vector_store_to_run:
litellm_params = vector_store_to_run.get("litellm_params", {}) or {}
data.update(litellm_params)
return data
@ -121,8 +96,7 @@ async def vector_store_search(
)
data = await _read_request_body(request=request)
if "vector_store_id" not in data:
data["vector_store_id"] = vector_store_id
data["vector_store_id"] = vector_store_id
# Check for legacy vector store registry (non-managed vector stores)
data = await _update_request_data_with_litellm_managed_vector_store_registry(

View file

@ -1,7 +1,9 @@
import json
from typing import Any, Dict, Literal, Optional
from fastapi import HTTPException, Request
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import (
LiteLLM_ObjectPermissionTable,
@ -13,6 +15,21 @@ from litellm.types.vector_stores import LiteLLM_ManagedVectorStore
from litellm.utils import ProviderConfigManager
def _normalize_litellm_params(
vector_store: LiteLLM_ManagedVectorStore,
) -> LiteLLM_ManagedVectorStore:
litellm_params = vector_store.get("litellm_params")
if isinstance(litellm_params, str):
normalized = LiteLLM_ManagedVectorStore(**dict(vector_store))
try:
parsed = json.loads(litellm_params)
normalized["litellm_params"] = parsed if isinstance(parsed, dict) else {}
except (TypeError, ValueError):
normalized["litellm_params"] = {}
return normalized
return vector_store
def _is_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> bool:
return (
user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
@ -120,6 +137,104 @@ async def can_user_access_vector_store(
return False
async def get_litellm_managed_vector_store(
vector_store_id: str,
) -> Optional[LiteLLM_ManagedVectorStore]:
"""
Resolve a LiteLLM-managed vector store from the registry or shared cache.
Provider-native vector store IDs will not be present in either location and
return None, preserving direct provider behavior while still protecting
LiteLLM-managed multi-tenant stores.
"""
if not vector_store_id:
return None
if litellm.vector_store_registry is not None:
try:
vector_store = litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry(
vector_store_id=vector_store_id
)
if vector_store is not None:
return _normalize_litellm_params(vector_store)
except Exception as e:
verbose_proxy_logger.warning(
"Failed to resolve vector store id=%s from registry: %s",
vector_store_id,
e,
)
raise HTTPException(
status_code=500,
detail="Unable to validate vector store access",
) from e
try:
from litellm.proxy.auth.auth_checks import (
get_managed_vector_store_rows_by_uuids,
)
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
if prisma_client is None:
return None
rows = await get_managed_vector_store_rows_by_uuids(
uuids=[vector_store_id],
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
if not rows:
return None
return _normalize_litellm_params(
LiteLLM_ManagedVectorStore(**rows[0].model_dump())
)
except Exception as e:
verbose_proxy_logger.warning(
"Failed to resolve vector store id=%s from shared cache: %s",
vector_store_id,
e,
)
raise HTTPException(
status_code=500,
detail="Unable to validate vector store access",
) from e
async def assert_user_can_access_vector_store(
vector_store: LiteLLM_ManagedVectorStore,
user_api_key_dict: UserAPIKeyAuth,
detail: str = "Access denied: You do not have permission to access this vector store",
) -> None:
"""Raise 403 unless the caller can access the resolved vector store."""
if not await can_user_access_vector_store(vector_store, user_api_key_dict):
raise HTTPException(status_code=403, detail=detail)
async def assert_user_can_access_vector_store_id(
vector_store_id: str,
user_api_key_dict: UserAPIKeyAuth,
detail: str = "Access denied: You do not have permission to access this vector store",
) -> Optional[LiteLLM_ManagedVectorStore]:
"""
Resolve a managed vector store id and enforce ownership if it exists.
Unknown ids are treated as provider-native ids and are not rejected here.
"""
vector_store = await get_litellm_managed_vector_store(
vector_store_id=vector_store_id
)
if vector_store is not None:
await assert_user_can_access_vector_store(
vector_store=vector_store,
user_api_key_dict=user_api_key_dict,
detail=detail,
)
return vector_store
def _does_endpoint_match(endpoint_path: str, request_path: str) -> bool:
if endpoint_path in request_path:
return True

View file

@ -17,9 +17,11 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
prepare_data_with_credentials,
)
from litellm.proxy.vector_store_endpoints.utils import (
assert_user_can_access_vector_store_id,
is_allowed_to_call_vector_store_files_endpoint,
)
from litellm.types.utils import LlmProviders
from litellm.types.vector_stores import LiteLLM_ManagedVectorStore
if TYPE_CHECKING:
from litellm.router import Router
@ -193,6 +195,8 @@ def _update_request_data_with_litellm_managed_vector_store_registry(
data: Dict,
vector_store_id: str,
llm_router: Optional["Router"] = None,
managed_vector_store: Optional[LiteLLM_ManagedVectorStore] = None,
should_lookup_registry: bool = True,
) -> Dict:
"""
Update request data with model routing information from managed vector store.
@ -262,23 +266,27 @@ def _update_request_data_with_litellm_managed_vector_store_registry(
return data
# Legacy path: Check vector store registry for non-managed vector stores
if litellm.vector_store_registry is not None:
# Legacy path: Check vector store registry for non-managed vector stores.
vector_store_to_run = managed_vector_store
if (
vector_store_to_run is None
and should_lookup_registry
and litellm.vector_store_registry is not None
):
vector_store_to_run = litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry(
vector_store_id=vector_store_id
)
if vector_store_to_run is not None:
if "custom_llm_provider" in vector_store_to_run:
data["custom_llm_provider"] = vector_store_to_run.get(
"custom_llm_provider"
)
if "litellm_credential_name" in vector_store_to_run:
data["litellm_credential_name"] = vector_store_to_run.get(
"litellm_credential_name"
)
if "litellm_params" in vector_store_to_run:
litellm_params = vector_store_to_run.get("litellm_params", {}) or {}
data.update(litellm_params)
if vector_store_to_run is not None:
if "custom_llm_provider" in vector_store_to_run:
data["custom_llm_provider"] = vector_store_to_run.get("custom_llm_provider")
if "litellm_credential_name" in vector_store_to_run:
data["litellm_credential_name"] = vector_store_to_run.get(
"litellm_credential_name"
)
if "litellm_params" in vector_store_to_run:
litellm_params = vector_store_to_run.get("litellm_params", {}) or {}
data.update(litellm_params)
return data
@ -363,8 +371,11 @@ async def vector_store_file_create(
)
data = await _read_request_body(request=request)
if "vector_store_id" not in data:
data["vector_store_id"] = vector_store_id
data["vector_store_id"] = vector_store_id
managed_vector_store = await assert_user_can_access_vector_store_id(
vector_store_id=vector_store_id,
user_api_key_dict=user_api_key_dict,
)
# Handle managed file IDs if present in request body
original_managed_file_id = None
@ -375,7 +386,11 @@ async def vector_store_file_create(
# Then handle managed vector store IDs
data = _update_request_data_with_litellm_managed_vector_store_registry(
data=data, vector_store_id=vector_store_id, llm_router=llm_router
data=data,
vector_store_id=vector_store_id,
llm_router=llm_router,
managed_vector_store=managed_vector_store,
should_lookup_registry=False,
)
provider_enum = await _resolve_provider(data=data, request=request)
@ -459,9 +474,18 @@ async def vector_store_file_list(
query_params = dict(request.query_params)
data: Dict[str, Optional[str]] = {"vector_store_id": vector_store_id}
data.update(query_params)
data["vector_store_id"] = vector_store_id
managed_vector_store = await assert_user_can_access_vector_store_id(
vector_store_id=vector_store_id,
user_api_key_dict=user_api_key_dict,
)
data = _update_request_data_with_litellm_managed_vector_store_registry(
data=data, vector_store_id=vector_store_id, llm_router=llm_router
data=data,
vector_store_id=vector_store_id,
llm_router=llm_router,
managed_vector_store=managed_vector_store,
should_lookup_registry=False,
)
provider_enum = await _resolve_provider(data=data, request=request)
@ -541,6 +565,10 @@ async def vector_store_file_retrieve(
"vector_store_id": vector_store_id,
"file_id": file_id,
}
managed_vector_store = await assert_user_can_access_vector_store_id(
vector_store_id=vector_store_id,
user_api_key_dict=user_api_key_dict,
)
# Handle managed file IDs first
data, original_managed_file_id = _update_request_data_with_managed_file_id(
@ -549,7 +577,11 @@ async def vector_store_file_retrieve(
# Then handle managed vector store IDs
data = _update_request_data_with_litellm_managed_vector_store_registry(
data=data, vector_store_id=vector_store_id, llm_router=llm_router
data=data,
vector_store_id=vector_store_id,
llm_router=llm_router,
managed_vector_store=managed_vector_store,
should_lookup_registry=False,
)
provider_enum = await _resolve_provider(data=data, request=request)
@ -635,6 +667,10 @@ async def vector_store_file_content(
"vector_store_id": vector_store_id,
"file_id": file_id,
}
managed_vector_store = await assert_user_can_access_vector_store_id(
vector_store_id=vector_store_id,
user_api_key_dict=user_api_key_dict,
)
# Handle managed file IDs first
data, original_managed_file_id = _update_request_data_with_managed_file_id(
@ -643,7 +679,11 @@ async def vector_store_file_content(
# Then handle managed vector store IDs
data = _update_request_data_with_litellm_managed_vector_store_registry(
data=data, vector_store_id=vector_store_id, llm_router=llm_router
data=data,
vector_store_id=vector_store_id,
llm_router=llm_router,
managed_vector_store=managed_vector_store,
should_lookup_registry=False,
)
provider_enum = await _resolve_provider(data=data, request=request)
@ -729,6 +769,10 @@ async def vector_store_file_update(
data = await _read_request_body(request=request)
data["vector_store_id"] = vector_store_id
data["file_id"] = file_id
managed_vector_store = await assert_user_can_access_vector_store_id(
vector_store_id=vector_store_id,
user_api_key_dict=user_api_key_dict,
)
# Handle managed file IDs first
data, original_managed_file_id = _update_request_data_with_managed_file_id(
@ -737,7 +781,11 @@ async def vector_store_file_update(
# Then handle managed vector store IDs
data = _update_request_data_with_litellm_managed_vector_store_registry(
data=data, vector_store_id=vector_store_id, llm_router=llm_router
data=data,
vector_store_id=vector_store_id,
llm_router=llm_router,
managed_vector_store=managed_vector_store,
should_lookup_registry=False,
)
provider_enum = await _resolve_provider(data=data, request=request)
@ -823,6 +871,10 @@ async def vector_store_file_delete(
"vector_store_id": vector_store_id,
"file_id": file_id,
}
managed_vector_store = await assert_user_can_access_vector_store_id(
vector_store_id=vector_store_id,
user_api_key_dict=user_api_key_dict,
)
# Handle managed file IDs first
data, original_managed_file_id = _update_request_data_with_managed_file_id(
@ -831,7 +883,11 @@ async def vector_store_file_delete(
# Then handle managed vector store IDs
data = _update_request_data_with_litellm_managed_vector_store_registry(
data=data, vector_store_id=vector_store_id, llm_router=llm_router
data=data,
vector_store_id=vector_store_id,
llm_router=llm_router,
managed_vector_store=managed_vector_store,
should_lookup_registry=False,
)
provider_enum = await _resolve_provider(data=data, request=request)

View file

@ -12,11 +12,13 @@ import base64
import os
from base64 import b64encode
from typing import Optional
from urllib.parse import unquote
import httpx
from fastapi import APIRouter, Request, Response
from fastapi import APIRouter, HTTPException, Request, Response, status
import litellm
from litellm.litellm_core_utils.url_utils import SSRFError, validate_url
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
@ -27,6 +29,7 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
router = APIRouter()
default_vertex_config = None
_DEFAULT_LANGFUSE_HOST = "https://cloud.langfuse.com"
def create_request_copy(request: Request):
@ -39,6 +42,116 @@ def create_request_copy(request: Request):
}
def _decode_to_convergence(value: str) -> str:
previous = value
while True:
decoded = unquote(previous)
if decoded == previous:
return decoded
previous = decoded
def _normalize_langfuse_base_url(base_target_url: str) -> str:
if not (
base_target_url.startswith("http://") or base_target_url.startswith("https://")
):
# Existing behavior allows host-only Langfuse settings.
base_target_url = "http://" + base_target_url
try:
base_url = httpx.URL(base_target_url)
except Exception as e:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={"error": f"Invalid Langfuse host: {str(e)}"},
)
if base_url.scheme not in ("http", "https") or not base_url.host:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={"error": "Invalid Langfuse host"},
)
if base_url.userinfo:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={"error": "Langfuse host must not include credentials"},
)
return str(base_url)
def _validate_langfuse_proxy_path(endpoint: str) -> str:
decoded_endpoint = _decode_to_convergence(endpoint)
if any(ord(char) < 32 for char in decoded_endpoint):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={"error": "Invalid Langfuse endpoint path"},
)
if "\\" in decoded_endpoint or decoded_endpoint.startswith("//"):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={"error": "Invalid Langfuse endpoint path"},
)
endpoint_path = "/" + decoded_endpoint.lstrip("/")
if any(segment in (".", "..") for segment in endpoint_path.split("/")):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={"error": "Invalid Langfuse endpoint path"},
)
return endpoint_path
def _get_langfuse_proxy_credentials(
*,
dynamic_host_supplied: bool,
dynamic_langfuse_public_key: Optional[str],
dynamic_langfuse_secret_key: Optional[str],
):
if dynamic_host_supplied:
if not dynamic_langfuse_public_key or not dynamic_langfuse_secret_key:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={
"error": "Dynamic Langfuse hosts must include dynamic Langfuse credentials"
},
)
return dynamic_langfuse_public_key, dynamic_langfuse_secret_key
return (
dynamic_langfuse_public_key
or litellm.utils.get_secret(secret_name="LANGFUSE_PUBLIC_KEY"),
dynamic_langfuse_secret_key
or litellm.utils.get_secret(secret_name="LANGFUSE_SECRET_KEY"),
)
def _build_langfuse_proxy_target(
*,
endpoint: str,
base_target_url: str,
dynamic_host_supplied: bool,
):
endpoint_path = _validate_langfuse_proxy_path(endpoint)
base_url = httpx.URL(_normalize_langfuse_base_url(base_target_url))
updated_url = base_url.copy_with(path=endpoint_path)
custom_headers = {}
if dynamic_host_supplied and getattr(litellm, "user_url_validation", True):
try:
target_url, host_header = validate_url(str(updated_url))
except SSRFError as e:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={"error": f"Invalid Langfuse host: {str(e)}"},
)
custom_headers["Host"] = host_header
return target_url, custom_headers
return str(updated_url), custom_headers
@router.api_route(
"/langfuse/{endpoint:path}",
methods=["GET", "POST", "PUT", "DELETE", "PATCH"],
@ -91,44 +204,33 @@ async def langfuse_proxy_route(
elif k == "langfuse_host":
dynamic_langfuse_host = v
dynamic_host_supplied = dynamic_langfuse_host is not None
base_target_url: str = (
dynamic_langfuse_host
or os.getenv("LANGFUSE_HOST", "https://cloud.langfuse.com")
or "https://cloud.langfuse.com"
or os.getenv("LANGFUSE_HOST", _DEFAULT_LANGFUSE_HOST)
or _DEFAULT_LANGFUSE_HOST
)
if not (
base_target_url.startswith("http://") or base_target_url.startswith("https://")
):
# add http:// if unset, assume communicating over private network - e.g. render
base_target_url = "http://" + base_target_url
encoded_endpoint = httpx.URL(endpoint).path
# Ensure endpoint starts with '/' for proper URL construction
if not encoded_endpoint.startswith("/"):
encoded_endpoint = "/" + encoded_endpoint
# Construct the full target URL using httpx
base_url = httpx.URL(base_target_url)
updated_url = base_url.copy_with(path=encoded_endpoint)
# Add or update query parameters
langfuse_public_key = dynamic_langfuse_public_key or litellm.utils.get_secret(
secret_name="LANGFUSE_PUBLIC_KEY"
langfuse_public_key, langfuse_secret_key = _get_langfuse_proxy_credentials(
dynamic_host_supplied=dynamic_host_supplied,
dynamic_langfuse_public_key=dynamic_langfuse_public_key,
dynamic_langfuse_secret_key=dynamic_langfuse_secret_key,
)
langfuse_secret_key = dynamic_langfuse_secret_key or litellm.utils.get_secret(
secret_name="LANGFUSE_SECRET_KEY"
target_url, target_headers = _build_langfuse_proxy_target(
endpoint=endpoint,
base_target_url=base_target_url,
dynamic_host_supplied=dynamic_host_supplied,
)
langfuse_combined_key = "Basic " + b64encode(
f"{langfuse_public_key}:{langfuse_secret_key}".encode("utf-8")
).decode("ascii")
target_headers["Authorization"] = langfuse_combined_key
## CREATE PASS-THROUGH
endpoint_func = create_pass_through_route(
endpoint=endpoint,
target=str(updated_url),
custom_headers={"Authorization": langfuse_combined_key},
target=target_url,
custom_headers=target_headers,
query_params=dict(request.query_params), # type: ignore
) # dynamically construct pass-through endpoint based on incoming path
received_value = await endpoint_func(

View file

@ -377,15 +377,11 @@ def search(
_is_async = kwargs.pop("asearch", False) is True
# pull credentials from registry if available
vector_store_id_for_credentials = kwargs.get("vector_store_id", vector_store_id)
if (
litellm.vector_store_registry is not None
and vector_store_id_for_credentials is not None
):
if litellm.vector_store_registry is not None and vector_store_id is not None:
try:
registry_credentials = (
litellm.vector_store_registry.get_credentials_for_vector_store(
vector_store_id_for_credentials
vector_store_id
)
)
kwargs.update(registry_credentials)

View file

@ -0,0 +1,129 @@
import sys
from types import ModuleType, SimpleNamespace
import litellm
from litellm.integrations.langfuse.langfuse import resolve_langfuse_credentials
from litellm.integrations.langfuse.langfuse_handler import LangFuseHandler
def test_resolve_langfuse_credentials_does_not_use_env_for_dynamic_host(monkeypatch):
monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "global-public")
monkeypatch.setenv("LANGFUSE_SECRET_KEY", "global-secret")
public_key, secret_key, host = resolve_langfuse_credentials(
langfuse_host="https://attacker.example",
allow_env_credentials=False,
)
assert public_key is None
assert secret_key is None
assert host == "https://attacker.example"
def test_resolve_langfuse_credentials_accepts_secret_key_alias_for_dynamic_host(
monkeypatch,
):
monkeypatch.setenv("LANGFUSE_SECRET_KEY", "global-secret")
public_key, secret_key, host = resolve_langfuse_credentials(
langfuse_public_key="dynamic-public",
langfuse_secret_key="dynamic-secret",
langfuse_host="https://team-langfuse.example",
allow_env_credentials=False,
)
assert public_key == "dynamic-public"
assert secret_key == "dynamic-secret"
assert host == "https://team-langfuse.example"
def test_resolve_langfuse_credentials_keeps_env_for_global_config(monkeypatch):
monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "global-public")
monkeypatch.setenv("LANGFUSE_SECRET_KEY", "global-secret")
public_key, secret_key, host = resolve_langfuse_credentials(
langfuse_host="https://admin-configured.example",
allow_env_credentials=True,
)
assert public_key == "global-public"
assert secret_key == "global-secret"
assert host == "https://admin-configured.example"
def test_upstream_langfuse_debug_env_is_passed(monkeypatch):
from litellm.integrations.langfuse.langfuse import LangFuseLogger
class FakeLangfuse:
instances = []
def __init__(self, **kwargs):
self.kwargs = kwargs
FakeLangfuse.instances.append(self)
fake_langfuse_module = ModuleType("langfuse")
fake_langfuse_module.Langfuse = FakeLangfuse
fake_langfuse_module.version = SimpleNamespace(__version__="2.6.0")
monkeypatch.setitem(sys.modules, "langfuse", fake_langfuse_module)
monkeypatch.setattr(litellm, "initialized_langfuse_clients", 0)
monkeypatch.setenv("LANGFUSE_MOCK", "true")
monkeypatch.setenv("UPSTREAM_LANGFUSE_SECRET_KEY", "upstream-secret")
monkeypatch.setenv("UPSTREAM_LANGFUSE_PUBLIC_KEY", "upstream-public")
monkeypatch.setenv("UPSTREAM_LANGFUSE_HOST", "https://upstream.example")
monkeypatch.setenv("UPSTREAM_LANGFUSE_RELEASE", "release")
monkeypatch.setenv("UPSTREAM_LANGFUSE_DEBUG", "true")
logger = LangFuseLogger(
langfuse_public_key="public",
langfuse_secret="secret",
langfuse_host="https://langfuse.example",
)
assert logger.upstream_langfuse_debug == "true"
assert FakeLangfuse.instances[-1].kwargs["debug"] is True
def test_langfuse_handler_accepts_secret_key_alias(monkeypatch):
captured = {}
class FakeLangFuseLogger:
def __init__(
self,
*,
langfuse_public_key=None,
langfuse_secret=None,
langfuse_host=None,
allow_env_credentials=True,
):
captured["langfuse_public_key"] = langfuse_public_key
captured["langfuse_secret"] = langfuse_secret
captured["langfuse_host"] = langfuse_host
captured["allow_env_credentials"] = allow_env_credentials
class FakeDynamicLoggingCache:
def set_cache(self, *, credentials, service_name, logging_obj):
captured["cached_credentials"] = credentials
captured["cached_service_name"] = service_name
captured["cached_logging_obj"] = logging_obj
monkeypatch.setattr(
"litellm.integrations.langfuse.langfuse_handler.LangFuseLogger",
FakeLangFuseLogger,
)
logger = LangFuseHandler._create_langfuse_logger_from_credentials(
credentials={
"langfuse_public_key": "dynamic-public",
"langfuse_secret_key": "dynamic-secret",
"langfuse_host": "https://langfuse.example",
},
in_memory_dynamic_logger_cache=FakeDynamicLoggingCache(),
)
assert captured["langfuse_public_key"] == "dynamic-public"
assert captured["langfuse_secret"] == "dynamic-secret"
assert captured["langfuse_host"] == "https://langfuse.example"
assert captured["allow_env_credentials"] is False
assert captured["cached_service_name"] == "langfuse"
assert captured["cached_logging_obj"] is logger

View file

@ -0,0 +1,50 @@
import pytest
from litellm.integrations.langsmith import LangsmithLogger
@pytest.mark.asyncio
async def test_get_credentials_from_env_does_not_use_env_for_dynamic_base_url(
monkeypatch,
):
monkeypatch.setenv("LANGSMITH_API_KEY", "global-key")
monkeypatch.setenv("LANGSMITH_PROJECT", "global-project")
monkeypatch.setenv("LANGSMITH_TENANT_ID", "global-tenant")
logger = LangsmithLogger(
langsmith_api_key="default-key",
langsmith_project="default-project",
langsmith_base_url="https://default.example",
)
credentials = logger.get_credentials_from_env(
langsmith_base_url="https://attacker.example",
allow_env_credentials=False,
)
assert credentials["LANGSMITH_API_KEY"] is None
assert credentials["LANGSMITH_PROJECT"] == "litellm-completion"
assert credentials["LANGSMITH_BASE_URL"] == "https://attacker.example"
assert credentials["LANGSMITH_TENANT_ID"] is None
@pytest.mark.asyncio
async def test_dynamic_langsmith_base_url_does_not_inherit_default_api_key(
monkeypatch,
):
monkeypatch.setenv("LANGSMITH_API_KEY", "global-key")
logger = LangsmithLogger(
langsmith_api_key="default-key",
langsmith_project="default-project",
langsmith_base_url="https://default.example",
)
credentials = logger._get_credentials_to_use_for_request(
kwargs={
"standard_callback_dynamic_params": {
"langsmith_base_url": "https://attacker.example"
}
}
)
assert credentials["LANGSMITH_API_KEY"] is None
assert credentials["LANGSMITH_BASE_URL"] == "https://attacker.example"

View file

@ -17,7 +17,10 @@ import litellm
from litellm.proxy._types import (
CallInfo,
Litellm_EntityType,
LiteLLM_BudgetTable,
LiteLLM_EndUserTable,
LiteLLM_ObjectPermissionTable,
LiteLLM_TagTable,
LiteLLM_TeamTable,
LiteLLM_UserTable,
LitellmUserRoles,
@ -29,10 +32,12 @@ from litellm.proxy._types import (
from litellm.proxy.auth.auth_checks import (
ExperimentalUIJWTToken,
_can_object_call_vector_stores,
_check_end_user_budget,
_check_team_member_budget,
_get_fuzzy_user_object,
_get_team_db_check,
_log_budget_lookup_failure,
_tag_max_budget_check,
_team_max_budget_check,
_virtual_key_max_budget_alert_check,
_virtual_key_max_budget_check,
@ -1964,6 +1969,67 @@ async def test_team_budget_check_reads_from_spend_counter():
assert exc_info.value.current_cost == 1.5
@pytest.mark.asyncio
async def test_end_user_budget_check_reads_from_spend_counter():
"""End-user budget check should use get_current_spend when counter exists."""
end_user_object = LiteLLM_EndUserTable(
user_id="customer-1",
blocked=False,
spend=0.0,
litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0),
)
async def mock_get_current_spend(counter_key, fallback_spend):
if counter_key == "spend:end_user:customer-1":
return 1.5
return fallback_spend
with patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend):
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await _check_end_user_budget(
end_user_obj=end_user_object,
route="/chat/completions",
)
assert exc_info.value.current_cost == 1.5
assert exc_info.value.max_budget == 1.0
@pytest.mark.asyncio
async def test_tag_budget_check_reads_from_spend_counter():
"""Tag budget check should use get_current_spend when counter exists."""
from litellm.proxy.utils import ProxyLogging
tag_object = LiteLLM_TagTable(
tag_name="paid-tag",
spend=0.0,
litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0),
)
async def mock_get_current_spend(counter_key, fallback_spend):
if counter_key == "spend:tag:paid-tag":
return 1.5
return fallback_spend
with (
patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend),
patch(
"litellm.proxy.auth.auth_checks.get_tag_objects_batch",
new_callable=AsyncMock,
return_value={"paid-tag": tag_object},
),
):
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await _tag_max_budget_check(
request_body={"metadata": {"tags": ["paid-tag"]}},
prisma_client=MagicMock(),
user_api_key_cache=MagicMock(),
proxy_logging_obj=ProxyLogging(user_api_key_cache=None),
valid_token=UserAPIKeyAuth(token="test-token"),
)
assert exc_info.value.current_cost == 1.5
assert exc_info.value.max_budget == 1.0
@pytest.mark.asyncio
async def test_team_member_budget_check_reads_from_spend_counter():
"""Team member budget check should use get_current_spend when counter exists."""

View file

@ -2,6 +2,7 @@
Unit tests for auth_utils functions related to rate limiting and customer ID extraction.
"""
import base64
from typing import Optional
from unittest.mock import MagicMock, patch
@ -10,11 +11,12 @@ import pytest
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_utils import (
_get_customer_id_from_standard_headers,
abbreviate_api_key,
check_complete_credentials,
get_end_user_id_from_request_body,
get_model_from_request,
get_key_model_rpm_limit,
get_key_model_tpm_limit,
get_model_from_request,
get_project_model_rpm_limit,
get_project_model_tpm_limit,
is_request_body_safe,
@ -258,6 +260,206 @@ def test_get_model_from_request_vertex_passthrough_still_works():
assert get_model_from_request(request_data={}, route=route) == "gemini-1.5-pro"
def test_get_model_from_request_openai_deployment_route_still_works():
assert (
get_model_from_request(
request_data={},
route="/openai/deployments/my-azure-deployment/chat/completions",
)
== "my-azure-deployment"
)
def test_get_model_from_request_includes_file_endpoint_header_model():
assert (
get_model_from_request(
request_data={},
route="/v1/files",
request_headers={"X-LiteLLM-Model": "restricted-model"},
)
== "restricted-model"
)
def test_get_model_from_request_ignores_routing_header_on_standard_llm_routes():
assert (
get_model_from_request(
request_data={"model": "allowed-model"},
route="/v1/chat/completions",
request_headers={"x-litellm-model": "restricted-model"},
)
== "allowed-model"
)
def test_get_model_from_request_authorizes_all_file_routing_model_sources():
models = get_model_from_request(
request_data={"model": "body-model"},
route="/v1/files",
request_headers={"x-litellm-model": "header-model"},
request_query_params={"target_model_names": "query-model-a,query-model-b"},
)
assert isinstance(models, list)
assert set(models) == {
"body-model",
"query-model-a",
"query-model-b",
"header-model",
}
def test_get_model_from_request_extracts_simple_encoded_file_id_model():
from litellm.proxy.openai_files_endpoints.common_utils import (
encode_file_id_with_model,
)
file_id = encode_file_id_with_model(
file_id="file-provider-id",
model="restricted-model",
)
assert (
get_model_from_request(
request_data={"file_id": file_id},
route="/v1/files/{file_id}",
)
== "restricted-model"
)
def test_get_model_from_request_extracts_unified_file_id_models():
raw_unified_file_id = (
"litellm_proxy:application/octet-stream;unified_id,test-id;"
"target_model_names,model-a,model-b;llm_output_file_id,file-provider-id"
)
encoded_unified_file_id = (
base64.urlsafe_b64encode(raw_unified_file_id.encode()).decode().rstrip("=")
)
assert get_model_from_request(
request_data={"file_id": encoded_unified_file_id},
route="/v1/files/{file_id}",
) == ["model-a", "model-b"]
def test_get_model_from_request_extracts_eval_completion_model():
assert (
get_model_from_request(
request_data={"completion": {"model": "judge-model"}},
route="/v1/evals/{eval_id}/runs",
)
== "judge-model"
)
def test_get_model_from_request_includes_fine_tuning_target_model_query():
assert (
get_model_from_request(
request_data={},
route="/v1/fine_tuning/jobs",
request_query_params={"target_model_names": "fine-tune-model"},
)
== "fine-tune-model"
)
def test_get_model_from_request_extracts_video_id_model():
from litellm.types.videos.utils import encode_video_id_with_provider
video_id = encode_video_id_with_provider(
video_id="video-provider-id",
provider="openai",
model_id="video-model",
)
assert (
get_model_from_request(
request_data={"video_id": video_id},
route="/v1/videos/{video_id}",
)
== "video-model"
)
def test_get_model_from_request_only_runs_media_decoders_for_matching_fields():
with (
patch(
"litellm.types.videos.utils.decode_video_id_with_provider",
return_value={"model_id": "video-model"},
) as video_decoder,
patch(
"litellm.types.videos.utils.decode_character_id_with_provider",
return_value={"model_id": "character-model"},
) as character_decoder,
):
assert (
get_model_from_request(
request_data={"file_id": "file-provider-id"},
route="/v1/files/{file_id}",
)
is None
)
video_decoder.assert_not_called()
character_decoder.assert_not_called()
assert (
get_model_from_request(
request_data={"video_id": "video-provider-id"},
route="/v1/videos/{video_id}",
)
== "video-model"
)
video_decoder.assert_called_once_with("video-provider-id")
character_decoder.assert_not_called()
video_decoder.reset_mock()
character_decoder.reset_mock()
assert (
get_model_from_request(
request_data={"character_id": "character-provider-id"},
route="/v1/videos/{character_id}",
)
== "character-model"
)
video_decoder.assert_not_called()
character_decoder.assert_called_once_with("character-provider-id")
def test_get_model_from_request_handles_managed_id_decoder_failures():
with (
patch(
"litellm.proxy.openai_files_endpoints.common_utils.decode_model_from_file_id",
side_effect=Exception("decode failed"),
),
patch(
"litellm.llms.base_llm.managed_resources.utils.parse_unified_id",
side_effect=Exception("parse failed"),
),
patch(
"litellm.types.videos.utils.decode_video_id_with_provider",
side_effect=Exception("video decode failed"),
),
):
assert (
get_model_from_request(
request_data={"file_id": "not-a-managed-resource-id"},
route="/v1/files/{file_id}",
)
is None
)
assert (
get_model_from_request(
request_data={"video_id": "not-a-managed-resource-id"},
route="/v1/videos/{video_id}",
)
is None
)
def test_abbreviate_api_key():
assert abbreviate_api_key("sk-test-1234") == "sk-...1234"
def test_get_customer_user_header_returns_none_when_no_customer_role():
from litellm.proxy.auth.auth_utils import get_customer_user_header_from_mapping

View file

@ -1,8 +1,7 @@
import asyncio
import json
import os
import sys
from typing import Tuple
from types import SimpleNamespace
from unittest.mock import ANY, AsyncMock, MagicMock, patch
sys.path.insert(
@ -15,6 +14,8 @@ import litellm.proxy.proxy_server
from litellm.caching.dual_cache import DualCache
from litellm.proxy._types import (
LiteLLM_JWTAuth,
LiteLLM_BudgetTable,
LiteLLM_EndUserTable,
LiteLLM_UserTable,
LitellmUserRoles,
ProxyErrorTypes,
@ -23,8 +24,10 @@ from litellm.proxy._types import (
JWTRoutingOverride,
)
from litellm.proxy.auth.handle_jwt import JWTHandler
from litellm.proxy.auth.auth_checks import get_key_object, _cache_key_object
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.auth.user_api_key_auth import (
_reserve_budget_after_common_checks,
_run_centralized_common_checks,
_run_post_custom_auth_checks,
get_api_key,
@ -32,6 +35,13 @@ from litellm.proxy.auth.user_api_key_auth import (
)
class _RoutingRequest:
def __init__(self, headers=None, query_params=None):
self.headers = headers or {}
self.query_params = query_params or {}
self.state = SimpleNamespace()
def test_get_api_key():
bearer_token = "Bearer sk-12345678"
api_key = "sk-12345678"
@ -49,6 +59,74 @@ def test_get_api_key():
) == (api_key, passed_in_key)
@pytest.mark.asyncio
async def test_should_clear_stale_budget_reservation_when_budget_checks_skip():
user_api_key_auth_obj = UserAPIKeyAuth(
token="test_token",
budget_reservation={
"reserved_cost": 0.5,
"entries": [{"counter_key": "spend:key:test_token"}],
},
)
await _reserve_budget_after_common_checks(
user_api_key_auth_obj=user_api_key_auth_obj,
request_data={"model": "free-model"},
route="/v1/chat/completions",
llm_router=None,
team_object=None,
user_object=None,
prisma_client=None,
user_api_key_cache=MagicMock(),
proxy_logging_obj=MagicMock(),
skip_budget_checks=True,
)
assert user_api_key_auth_obj.budget_reservation is None
@pytest.mark.asyncio
async def test_should_not_reuse_cached_key_object_for_request_state():
key_cache = DualCache()
cached_key = UserAPIKeyAuth(
token="cached-token",
request_route="/old-route",
budget_reservation={
"reserved_cost": 0.5,
"entries": [{"counter_key": "spend:key:cached-token"}],
},
)
await _cache_key_object(
hashed_token="cached-token",
user_api_key_obj=cached_key,
user_api_key_cache=key_cache,
proxy_logging_obj=None,
)
first_request_key = await get_key_object(
hashed_token="cached-token",
prisma_client=MagicMock(),
user_api_key_cache=key_cache,
)
first_request_key.budget_reservation = {
"reserved_cost": 0.9,
"entries": [{"counter_key": "spend:key:cached-token"}],
}
first_request_key.request_route = "/chat/completions"
second_request_key = await get_key_object(
hashed_token="cached-token",
prisma_client=MagicMock(),
user_api_key_cache=key_cache,
)
assert first_request_key is not cached_key
assert second_request_key is not first_request_key
assert second_request_key.budget_reservation is None
assert second_request_key.request_route is None
@pytest.mark.asyncio
async def test_custom_auth_does_not_enforce_key_model_access_by_default():
valid_token = UserAPIKeyAuth(token="test_token", models=["gpt-4o-mini"])
@ -107,6 +185,39 @@ async def test_custom_auth_honors_key_level_model_access_restriction_allowed_wit
)
@pytest.mark.asyncio
async def test_custom_auth_enforces_key_model_access_from_file_route_header_with_opt_in():
valid_token = UserAPIKeyAuth(token="test_token", models=["allowed-model"])
request = _RoutingRequest(headers={"x-litellm-model": "restricted-model"})
with (
patch(
"litellm.proxy.auth.user_api_key_auth.can_key_call_model",
new_callable=AsyncMock,
) as mock_can_key,
patch(
"litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock
),
patch(
"litellm.proxy.proxy_server.general_settings",
{"custom_auth_run_common_checks": True},
),
):
await _run_post_custom_auth_checks(
valid_token=valid_token,
request=request,
request_data={},
route="/v1/files",
parent_otel_span=None,
)
mock_can_key.assert_awaited_once_with(
model="restricted-model",
llm_model_list=ANY,
valid_token=valid_token,
llm_router=ANY,
)
@pytest.mark.asyncio
async def test_custom_auth_honors_key_level_model_access_restriction_denied_with_opt_in():
valid_token = UserAPIKeyAuth(token="test_token", models=["gpt-4o-mini"])
@ -1752,7 +1863,11 @@ async def test_team_metadata_refreshed_from_team_object_during_auth():
from starlette.datastructures import URL
from starlette.requests import Request
from litellm.proxy._types import LiteLLM_TeamTableCachedObj, LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy._types import (
LiteLLM_TeamTableCachedObj,
LitellmUserRoles,
UserAPIKeyAuth,
)
from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder
api_key = "sk-test-team-metadata-refresh"
@ -1833,16 +1948,17 @@ async def test_team_metadata_refreshed_from_team_object_during_auth():
request_data={},
)
assert result.team_metadata == {"guardrails": ["test-guardrail-333"]}, (
f"team_metadata was not updated from fresh team object. Got: {result.team_metadata}"
)
assert result.team_metadata == {
"guardrails": ["test-guardrail-333"]
}, f"team_metadata was not updated from fresh team object. Got: {result.team_metadata}"
finally:
for k, v in _originals.items():
setattr(_proxy_server_mod, k, v)
# ---------------------------------------------------------------------------
# _run_centralized_common_checks — centralized authz gate
# ---------------------------------------------------------------------------
@ -1859,7 +1975,7 @@ def _proxy_attrs_for_centralized_checks(
"""
return {
"prisma_client": None,
"user_api_key_cache": MagicMock(),
"user_api_key_cache": DualCache(),
"proxy_logging_obj": MagicMock(),
"general_settings": ({"custom_auth_run_common_checks": True} if flag else {}),
"llm_router": None,
@ -2120,6 +2236,81 @@ async def test_centralized_common_checks_propagates_end_user_budget_error():
setattr(_proxy_server_mod, k, v)
@pytest.mark.asyncio
async def test_centralized_common_checks_reserves_request_end_user_budget():
"""Regression: reservation runs before user_api_key_auth() copies the
request end-user onto the token, so centralized checks must pass the
locally extracted end_user_id/end_user_object into reservation."""
import litellm.proxy.proxy_server as _proxy_server_mod
from fastapi import Request
from starlette.datastructures import URL
token = UserAPIKeyAuth(api_key="sk-test", user_id="u")
request = Request(scope={"type": "http", "headers": []})
request._url = URL(url="/chat/completions")
request_data = {
"model": "gpt-4o",
"messages": [{"role": "user", "content": "hello"}],
"user": "alice",
}
end_user_object = LiteLLM_EndUserTable(
user_id="alice",
blocked=False,
spend=0.0,
litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0),
)
attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=None)
counter_cache = DualCache()
attrs["spend_counter_cache"] = counter_cache
originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
try:
for k, v in attrs.items():
setattr(_proxy_server_mod, k, v)
with (
patch(
"litellm.proxy.auth.user_api_key_auth.get_end_user_object",
new_callable=AsyncMock,
return_value=end_user_object,
),
patch(
"litellm.proxy.auth.user_api_key_auth.common_checks",
new_callable=AsyncMock,
),
patch(
"litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost",
return_value=0.6,
),
):
assert token.end_user_id is None
await _run_centralized_common_checks(
user_api_key_auth_obj=token,
request=request,
request_data=request_data,
route="/chat/completions",
)
finally:
for k, v in originals.items():
setattr(_proxy_server_mod, k, v)
assert token.end_user_id is None
assert token.budget_reservation is not None
assert token.budget_reservation["entries"] == [
{
"counter_key": "spend:end_user:alice",
"entity_type": "EndUser",
"entity_id": "alice",
"reserved_cost": 0.6,
"applied_adjustment": 0.0,
}
]
assert counter_cache.in_memory_cache.get_cache(
key="spend:end_user:alice"
) == pytest.approx(0.6)
@pytest.mark.asyncio
async def test_centralized_common_checks_short_circuits_when_master_key_unset():
"""master_key=None is no-auth dev mode — admin-only routes and

View file

@ -0,0 +1,208 @@
"""
Unit tests for the personal-budget pre-call hook.
The reservation path (added in PR #26845) atomically pre-fills the same
`spend:user:{user_id}` counter this hook reads, admitting at a strict-`<`
boundary. Re-checking with `>=` after reservation would reject requests the
reservation already admitted when the reservation fills the counter to
exactly `max_budget` (e.g. requests with no `max_tokens` cap fall back to
reserving the smallest remaining headroom).
These tests pin the skip-when-reserved behavior and guard against drift.
"""
from unittest.mock import AsyncMock, patch
import pytest
from fastapi import HTTPException
from litellm.caching.caching import DualCache
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.hooks.max_budget_limiter import _PROXY_MaxBudgetLimiter
def _make_user_api_key_auth(
user_id: str = "user-1",
user_max_budget: float = 10.0,
user_spend: float = 0.0,
team_id=None,
budget_reservation=None,
) -> UserAPIKeyAuth:
return UserAPIKeyAuth(
api_key="sk-test",
user_id=user_id,
user_max_budget=user_max_budget,
user_spend=user_spend,
team_id=team_id,
budget_reservation=budget_reservation,
)
@pytest.mark.asyncio
async def test_under_budget_passes():
handler = _PROXY_MaxBudgetLimiter()
user_api_key_dict = _make_user_api_key_auth(user_max_budget=10.0)
with patch(
"litellm.proxy.proxy_server.get_current_spend",
new=AsyncMock(return_value=3.0),
):
result = await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=DualCache(),
data={},
call_type="completion",
)
assert result is None
@pytest.mark.asyncio
async def test_over_budget_rejects_without_reservation():
handler = _PROXY_MaxBudgetLimiter()
user_api_key_dict = _make_user_api_key_auth(user_max_budget=10.0)
with patch(
"litellm.proxy.proxy_server.get_current_spend",
new=AsyncMock(return_value=10.0),
):
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=DualCache(),
data={},
call_type="completion",
)
assert exc_info.value.status_code == 429
assert "Max budget limit reached." in exc_info.value.detail
@pytest.mark.asyncio
async def test_skips_when_user_counter_is_reserved():
"""
Reservation atomically pre-fills `spend:user:{user_id}` and admits the
request. The legacy `>=` check must not double-enforce on the same
counter — that's what produced the boundary regression where a fresh
user with no `max_tokens` cap got 429'd on their first request.
"""
handler = _PROXY_MaxBudgetLimiter()
user_api_key_dict = _make_user_api_key_auth(
user_id="user-1",
user_max_budget=10.0,
budget_reservation={
"reserved_cost": 10.0,
"entries": [
{
"counter_key": "spend:user:user-1",
"entity_type": "User",
"entity_id": "user-1",
"reserved_cost": 10.0,
"applied_adjustment": 0.0,
}
],
"finalized": False,
},
)
# `get_current_spend` would return 10.0 here (counter pre-filled by the
# reservation). The hook must skip without reading it.
with patch(
"litellm.proxy.proxy_server.get_current_spend",
new=AsyncMock(return_value=10.0),
) as mock_get_spend:
result = await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=DualCache(),
data={},
call_type="completion",
)
assert result is None
mock_get_spend.assert_not_awaited()
@pytest.mark.asyncio
async def test_does_not_skip_when_reservation_covers_a_different_counter():
"""
A reservation that only covers e.g. `spend:team:{team_id}` (not the user
counter) must not exempt the user-budget check.
"""
handler = _PROXY_MaxBudgetLimiter()
user_api_key_dict = _make_user_api_key_auth(
user_id="user-1",
user_max_budget=10.0,
budget_reservation={
"reserved_cost": 5.0,
"entries": [
{
"counter_key": "spend:team:team-x",
"entity_type": "Team",
"entity_id": "team-x",
"reserved_cost": 5.0,
"applied_adjustment": 0.0,
}
],
"finalized": False,
},
)
with patch(
"litellm.proxy.proxy_server.get_current_spend",
new=AsyncMock(return_value=10.0),
):
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=DualCache(),
data={},
call_type="completion",
)
assert exc_info.value.status_code == 429
@pytest.mark.asyncio
async def test_team_keys_skip_personal_budget():
handler = _PROXY_MaxBudgetLimiter()
user_api_key_dict = _make_user_api_key_auth(
user_max_budget=10.0,
team_id="team-1",
)
with patch(
"litellm.proxy.proxy_server.get_current_spend",
new=AsyncMock(return_value=999.0),
) as mock_get_spend:
result = await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=DualCache(),
data={},
call_type="completion",
)
assert result is None
mock_get_spend.assert_not_awaited()
@pytest.mark.asyncio
async def test_no_max_budget_passes():
handler = _PROXY_MaxBudgetLimiter()
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-test",
user_id="user-1",
)
with patch(
"litellm.proxy.proxy_server.get_current_spend",
new=AsyncMock(return_value=999.0),
) as mock_get_spend:
result = await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=DualCache(),
data={},
call_type="completion",
)
assert result is None
mock_get_spend.assert_not_awaited()

View file

@ -1,9 +1,7 @@
import json
import os
import sys
import pytest
from fastapi.testclient import TestClient
sys.path.insert(
0, os.path.abspath("../../../..")
@ -13,8 +11,11 @@ from datetime import datetime
from unittest.mock import AsyncMock, MagicMock, patch
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger
from litellm.types.utils import StandardLoggingPayload
from litellm.proxy.hooks.proxy_track_cost_callback import (
_ProxyDBLogger,
_get_budget_reservation_from_metadata,
_update_database_and_spend_counters,
)
@pytest.mark.asyncio
@ -62,7 +63,6 @@ async def test_async_post_call_failure_hook():
# Check the arguments passed to update_database
call_args = mock_update_database.call_args[1]
print("call_args", json.dumps(call_args, indent=4, default=str))
assert call_args["token"] == "test_api_key"
assert call_args["response_cost"] == 0.0
assert call_args["user_id"] == "test_user_id"
@ -128,6 +128,440 @@ async def test_async_post_call_failure_hook_non_llm_route():
mock_update_database.assert_not_called()
@pytest.mark.asyncio
async def test_async_post_call_failure_hook_releases_budget_reservation_before_route_skip():
logger = _ProxyDBLogger()
budget_reservation = {"reserved_cost": 0.5, "entries": []}
user_api_key_dict = UserAPIKeyAuth(
api_key="test_api_key",
request_route="/custom/route",
budget_reservation=budget_reservation,
)
with (
patch(
"litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation",
new_callable=AsyncMock,
) as mock_release_budget_reservation,
patch(
"litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database",
new_callable=AsyncMock,
) as mock_update_database,
):
await logger.async_post_call_failure_hook(
request_data={},
original_exception=Exception("Test exception"),
user_api_key_dict=user_api_key_dict,
)
assert mock_release_budget_reservation.await_count == 1
assert (
mock_release_budget_reservation.await_args.kwargs["budget_reservation"]
is user_api_key_dict.budget_reservation
)
mock_update_database.assert_not_called()
@pytest.mark.asyncio
async def test_should_continue_failure_tracking_when_budget_release_fails():
logger = _ProxyDBLogger()
budget_reservation = {"reserved_cost": 0.5, "entries": []}
user_api_key_dict = UserAPIKeyAuth(
api_key="test_api_key",
user_id="test_user_id",
team_id="test_team_id",
request_route="/chat/completions",
budget_reservation=budget_reservation,
)
with (
patch(
"litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation",
new_callable=AsyncMock,
side_effect=RuntimeError("redis unavailable"),
) as mock_release_budget_reservation,
patch(
"litellm.proxy.hooks.proxy_track_cost_callback._invalidate_budget_reservation_counters",
new_callable=AsyncMock,
) as mock_invalidate_budget_reservation_counters,
patch(
"litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database",
new_callable=AsyncMock,
) as mock_update_database,
patch(
"litellm.proxy.hooks.proxy_track_cost_callback.verbose_proxy_logger.exception",
) as mock_log_exception,
):
await logger.async_post_call_failure_hook(
request_data={
"model": "gpt-4",
"messages": [{"role": "user", "content": "Hello"}],
},
original_exception=Exception("provider failed"),
user_api_key_dict=user_api_key_dict,
)
assert mock_release_budget_reservation.await_count == 1
assert (
mock_release_budget_reservation.await_args.kwargs["budget_reservation"]
is user_api_key_dict.budget_reservation
)
assert mock_invalidate_budget_reservation_counters.await_count == 1
assert (
mock_invalidate_budget_reservation_counters.await_args.kwargs[
"budget_reservation"
]
is user_api_key_dict.budget_reservation
)
assert user_api_key_dict.budget_reservation["finalized"] is True
mock_log_exception.assert_called_once()
mock_update_database.assert_called_once()
@pytest.mark.asyncio
async def test_track_cost_callback_releases_budget_reservation_when_spend_tracking_skips():
logger = _ProxyDBLogger()
budget_reservation = {"reserved_cost": 0.5, "entries": []}
user_api_key_auth = UserAPIKeyAuth(budget_reservation=budget_reservation)
kwargs = {
"model": "gpt-4",
"litellm_params": {
"metadata": {
"user_api_key_auth": user_api_key_auth,
},
},
"standard_logging_object": {
"response_cost": 0.1,
"request_tags": None,
},
"stream": False,
}
with patch(
"litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation",
new_callable=AsyncMock,
) as mock_release_budget_reservation:
await logger._PROXY_track_cost_callback(
kwargs=kwargs,
completion_response=None,
start_time=datetime.now(),
end_time=datetime.now(),
)
mock_release_budget_reservation.assert_awaited_once_with(
budget_reservation=budget_reservation,
)
@pytest.mark.asyncio
async def test_track_cost_callback_releases_budget_reservation_when_response_cost_missing():
logger = _ProxyDBLogger()
budget_reservation = {"reserved_cost": 0.5, "entries": []}
user_api_key_auth = UserAPIKeyAuth(budget_reservation=budget_reservation)
kwargs = {
"model": "gpt-4",
"call_type": "acompletion",
"litellm_params": {
"metadata": {
"user_api_key_auth": user_api_key_auth,
},
},
"standard_logging_object": {
"response_cost": None,
"response_cost_failure_debug_info": "missing custom price",
"request_tags": None,
},
"stream": False,
}
with (
patch(
"litellm.proxy.proxy_server.proxy_logging_obj",
) as mock_proxy_logging,
patch(
"litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation",
new_callable=AsyncMock,
) as mock_release_budget_reservation,
):
mock_proxy_logging.failed_tracking_alert = AsyncMock()
await logger._PROXY_track_cost_callback(
kwargs=kwargs,
completion_response=None,
start_time=datetime.now(),
end_time=datetime.now(),
)
mock_release_budget_reservation.assert_awaited_once_with(
budget_reservation=budget_reservation,
)
def test_get_budget_reservation_from_metadata_handles_dict_auth_object():
budget_reservation = {
"reserved_cost": 0.5,
"entries": [{"counter_key": "spend:key:test_api_key"}],
}
assert (
_get_budget_reservation_from_metadata(
metadata={"user_api_key_auth": dict(UserAPIKeyAuth())}
)
is None
)
assert (
_get_budget_reservation_from_metadata(
metadata={
"user_api_key_auth": UserAPIKeyAuth(
budget_reservation=budget_reservation
)
}
)
== budget_reservation
)
assert (
_get_budget_reservation_from_metadata(
metadata={
"user_api_key_auth": dict(
UserAPIKeyAuth(budget_reservation=budget_reservation)
)
}
)
== budget_reservation
)
assert (
_get_budget_reservation_from_metadata(
metadata={"user_api_key_budget_reservation": budget_reservation}
)
is budget_reservation
)
@pytest.mark.asyncio
async def test_update_database_and_spend_counters_releases_reservation_when_db_update_fails():
proxy_logging_obj = MagicMock()
proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock(
side_effect=Exception("db unavailable")
)
increment_spend_counters = AsyncMock()
budget_reservation = {"reserved_cost": 0.5, "entries": []}
with patch(
"litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation",
new_callable=AsyncMock,
) as mock_release_budget_reservation:
with pytest.raises(Exception, match="db unavailable"):
await _update_database_and_spend_counters(
proxy_logging_obj=proxy_logging_obj,
increment_spend_counters=increment_spend_counters,
user_api_key="test_api_key",
user_id="test_user_id",
end_user_id=None,
team_id="test_team_id",
org_id="test_org_id",
kwargs={},
completion_response=None,
start_time=datetime.now(),
end_time=datetime.now(),
response_cost=0.2,
budget_reservation=budget_reservation,
)
mock_release_budget_reservation.assert_awaited_once_with(
budget_reservation=budget_reservation,
)
increment_spend_counters.assert_not_awaited()
@pytest.mark.asyncio
async def test_update_database_and_spend_counters_preserves_db_exception_when_release_fails():
proxy_logging_obj = MagicMock()
db_exception = RuntimeError("db unavailable")
proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock(
side_effect=db_exception
)
increment_spend_counters = AsyncMock()
budget_reservation = {"reserved_cost": 0.5, "entries": []}
with (
patch(
"litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation",
new_callable=AsyncMock,
side_effect=RuntimeError("release unavailable"),
) as mock_release_budget_reservation,
patch(
"litellm.proxy.hooks.proxy_track_cost_callback.verbose_proxy_logger.exception",
) as mock_log_exception,
patch(
"litellm.proxy.hooks.proxy_track_cost_callback._invalidate_budget_reservation_counters",
new_callable=AsyncMock,
side_effect=RuntimeError("invalidate unavailable"),
) as mock_invalidate_budget_reservation_counters,
):
with pytest.raises(RuntimeError) as exc_info:
await _update_database_and_spend_counters(
proxy_logging_obj=proxy_logging_obj,
increment_spend_counters=increment_spend_counters,
user_api_key="test_api_key",
user_id="test_user_id",
end_user_id=None,
team_id="test_team_id",
org_id="test_org_id",
kwargs={},
completion_response=None,
start_time=datetime.now(),
end_time=datetime.now(),
response_cost=0.2,
budget_reservation=budget_reservation,
)
assert exc_info.value is db_exception
mock_release_budget_reservation.assert_awaited_once_with(
budget_reservation=budget_reservation,
)
mock_invalidate_budget_reservation_counters.assert_awaited_once_with(
budget_reservation=budget_reservation,
)
assert mock_log_exception.call_count == 2
mock_log_exception.assert_any_call(
"Failed to release budget reservation after database update failed"
)
mock_log_exception.assert_any_call(
"Failed to invalidate budget reservation counters after release failed"
)
increment_spend_counters.assert_not_awaited()
@pytest.mark.asyncio
async def test_update_database_and_spend_counters_updates_counters_after_db_update():
proxy_logging_obj = MagicMock()
proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock()
increment_spend_counters = AsyncMock()
budget_reservation = {"reserved_cost": 0.5, "entries": []}
await _update_database_and_spend_counters(
proxy_logging_obj=proxy_logging_obj,
increment_spend_counters=increment_spend_counters,
user_api_key="test_api_key",
user_id="test_user_id",
end_user_id="test_end_user_id",
team_id="test_team_id",
org_id="test_org_id",
kwargs={},
completion_response=None,
start_time=datetime.now(),
end_time=datetime.now(),
response_cost=0.2,
budget_reservation=budget_reservation,
request_tags=["tag-a"],
)
proxy_logging_obj.db_spend_update_writer.update_database.assert_awaited_once()
increment_spend_counters.assert_awaited_once_with(
token="test_api_key",
team_id="test_team_id",
user_id="test_user_id",
response_cost=0.2,
org_id="test_org_id",
budget_reservation=budget_reservation,
end_user_id="test_end_user_id",
tags=["tag-a"],
)
@pytest.mark.asyncio
async def test_update_database_and_spend_counters_invalidates_reservation_when_counter_update_fails():
proxy_logging_obj = MagicMock()
proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock()
increment_spend_counters = AsyncMock(side_effect=Exception("counter unavailable"))
budget_reservation = {
"reserved_cost": 0.5,
"entries": [{"counter_key": "spend:key:test_api_key"}],
}
with patch(
"litellm.proxy.spend_tracking.budget_reservation.invalidate_budget_reservation_counters",
new_callable=AsyncMock,
) as mock_invalidate_budget_reservation_counters:
with pytest.raises(Exception, match="counter unavailable"):
await _update_database_and_spend_counters(
proxy_logging_obj=proxy_logging_obj,
increment_spend_counters=increment_spend_counters,
user_api_key="test_api_key",
user_id="test_user_id",
end_user_id=None,
team_id="test_team_id",
org_id="test_org_id",
kwargs={},
completion_response=None,
start_time=datetime.now(),
end_time=datetime.now(),
response_cost=0.2,
budget_reservation=budget_reservation,
)
mock_invalidate_budget_reservation_counters.assert_awaited_once_with(
budget_reservation=budget_reservation,
)
assert budget_reservation["finalized"] is True
proxy_logging_obj.db_spend_update_writer.update_database.assert_awaited_once()
@pytest.mark.asyncio
async def test_update_database_and_spend_counters_preserves_counter_exception_when_invalidation_fails():
proxy_logging_obj = MagicMock()
proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock()
counter_exception = RuntimeError("counter unavailable")
increment_spend_counters = AsyncMock(side_effect=counter_exception)
budget_reservation = {
"reserved_cost": 0.5,
"entries": [{"counter_key": "spend:key:test_api_key"}],
}
with (
patch(
"litellm.proxy.spend_tracking.budget_reservation.invalidate_budget_reservation_counters",
new_callable=AsyncMock,
side_effect=RuntimeError("invalidate unavailable"),
) as mock_invalidate_budget_reservation_counters,
patch(
"litellm.proxy.hooks.proxy_track_cost_callback.verbose_proxy_logger.exception",
) as mock_log_exception,
):
with pytest.raises(RuntimeError) as exc_info:
await _update_database_and_spend_counters(
proxy_logging_obj=proxy_logging_obj,
increment_spend_counters=increment_spend_counters,
user_api_key="test_api_key",
user_id="test_user_id",
end_user_id=None,
team_id="test_team_id",
org_id="test_org_id",
kwargs={},
completion_response=None,
start_time=datetime.now(),
end_time=datetime.now(),
response_cost=0.2,
budget_reservation=budget_reservation,
)
assert exc_info.value is counter_exception
mock_invalidate_budget_reservation_counters.assert_awaited_once_with(
budget_reservation=budget_reservation,
)
mock_log_exception.assert_called_once_with(
"Failed to invalidate budget reservation counters after spend counter update failed"
)
assert budget_reservation["finalized"] is True
proxy_logging_obj.db_spend_update_writer.update_database.assert_awaited_once()
@pytest.mark.asyncio
async def test_track_cost_callback_skips_when_no_standard_logging_object():
"""
@ -344,7 +778,7 @@ async def test_enrich_failure_metadata_skips_when_no_api_key():
"user_api_key_team_id": None,
"user_api_key_team_alias": None,
}
result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata)
await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata)
mock_get_key.assert_not_called()

File diff suppressed because it is too large Load diff

View file

@ -0,0 +1,102 @@
import socket
import pytest
from fastapi import HTTPException
import litellm
from litellm.proxy.vertex_ai_endpoints.langfuse_endpoints import (
_build_langfuse_proxy_target,
_get_langfuse_proxy_credentials,
)
def test_dynamic_langfuse_host_requires_dynamic_credentials(monkeypatch):
monkeypatch.setattr(litellm, "user_url_validation", True, raising=False)
monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "global-public")
monkeypatch.setenv("LANGFUSE_SECRET_KEY", "global-secret")
with pytest.raises(HTTPException) as exc:
_get_langfuse_proxy_credentials(
dynamic_host_supplied=True,
dynamic_langfuse_public_key=None,
dynamic_langfuse_secret_key=None,
)
assert exc.value.status_code == 400
def test_global_langfuse_host_can_use_env_credentials(monkeypatch):
monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "global-public")
monkeypatch.setenv("LANGFUSE_SECRET_KEY", "global-secret")
public_key, secret_key = _get_langfuse_proxy_credentials(
dynamic_host_supplied=False,
dynamic_langfuse_public_key=None,
dynamic_langfuse_secret_key=None,
)
assert public_key == "global-public"
assert secret_key == "global-secret"
@pytest.mark.parametrize(
"endpoint",
[
"../api/public/projects",
"%2e%2e/api/public/projects",
"%252e%252e%252fapi/public/projects",
"api\\public\\projects",
"%2f%2fattacker.example/api",
],
)
def test_langfuse_proxy_target_rejects_traversal_paths(endpoint):
with pytest.raises(HTTPException) as exc:
_build_langfuse_proxy_target(
endpoint=endpoint,
base_target_url="https://cloud.langfuse.com",
dynamic_host_supplied=False,
)
assert exc.value.status_code == 400
def test_dynamic_langfuse_proxy_target_rejects_internal_host(monkeypatch):
monkeypatch.setattr(litellm, "user_url_validation", True, raising=False)
with pytest.raises(HTTPException) as exc:
_build_langfuse_proxy_target(
endpoint="api/public/projects",
base_target_url="http://127.0.0.1:3000",
dynamic_host_supplied=True,
)
assert exc.value.status_code == 400
def test_dynamic_langfuse_proxy_target_preserves_host_header_for_http(monkeypatch):
monkeypatch.setattr(litellm, "user_url_validation", True, raising=False)
def fake_getaddrinfo(host, port, proto):
assert host == "langfuse.example"
assert port == 80
assert proto == socket.IPPROTO_TCP
return [
(
socket.AF_INET,
socket.SOCK_STREAM,
socket.IPPROTO_TCP,
"",
("8.8.8.8", 80),
)
]
monkeypatch.setattr(socket, "getaddrinfo", fake_getaddrinfo)
target_url, headers = _build_langfuse_proxy_target(
endpoint="api/public/projects",
base_target_url="http://langfuse.example",
dynamic_host_supplied=True,
)
assert target_url == "http://8.8.8.8/api/public/projects"
assert headers["Host"] == "langfuse.example"

View file

@ -1,6 +1,89 @@
import sys
from types import ModuleType, SimpleNamespace
from litellm.proxy._lazy_openapi_snapshot import _normalize_operation_ids
def test_generate_snapshot_uses_shared_operation_id_reservations(monkeypatch):
from litellm.proxy import _lazy_openapi_snapshot
route_a = SimpleNamespace(path="/feature-a/items")
route_b = SimpleNamespace(path="/feature-b/items")
fake_app = SimpleNamespace(
title="LiteLLM test",
version="0.0.0",
routes=[route_a, route_b],
)
fake_feature_a_module = ModuleType("fake_feature_a")
fake_feature_b_module = ModuleType("fake_feature_b")
monkeypatch.setitem(sys.modules, "fake_feature_a", fake_feature_a_module)
monkeypatch.setitem(sys.modules, "fake_feature_b", fake_feature_b_module)
fake_lazy_features_module = ModuleType("litellm.proxy._lazy_features")
fake_lazy_features_module.LAZY_FEATURES = [
SimpleNamespace(
name="feature-a",
module_path="fake_feature_a",
path_prefixes=("/feature-a",),
register_fn=lambda app, module: None,
),
SimpleNamespace(
name="feature-b",
module_path="fake_feature_b",
path_prefixes=("/feature-b",),
register_fn=lambda app, module: None,
),
]
monkeypatch.setitem(
sys.modules, "litellm.proxy._lazy_features", fake_lazy_features_module
)
def fake_get_openapi(title, version, routes):
path = routes[0].path
return {
"paths": {path: {"get": {"operationId": "shared_operation_id_get"}}},
"components": {"schemas": {"Example": {"type": "object"}}},
}
def fake_ensure_unique_openapi_operation_ids(schema, reserved_operation_ids):
for path_item in schema["paths"].values():
operation = path_item["get"]
operation_id = operation["operationId"]
if operation_id in reserved_operation_ids:
operation_id = f"{operation_id}_2"
operation["operationId"] = operation_id
reserved_operation_ids.add(operation_id)
return schema
fake_proxy_server_module = ModuleType("litellm.proxy.proxy_server")
fake_proxy_server_module.app = fake_app
fake_proxy_server_module.ensure_unique_openapi_operation_ids = (
fake_ensure_unique_openapi_operation_ids
)
monkeypatch.setitem(
sys.modules, "litellm.proxy.proxy_server", fake_proxy_server_module
)
monkeypatch.setattr("fastapi.openapi.utils.get_openapi", fake_get_openapi)
fragments = _lazy_openapi_snapshot.generate_snapshot()
assert (
fragments["feature-a"]["paths"]["/feature-a/items"]["get"]["operationId"]
== "shared_operation_id_get"
)
assert (
fragments["feature-b"]["paths"]["/feature-b/items"]["get"]["operationId"]
== "shared_operation_id_get_2"
)
assert fragments["feature-a"]["paths"]["/feature-a/items"]["get"]["tags"] == [
"feature-a"
]
assert fragments["feature-b"]["paths"]["/feature-b/items"]["get"]["tags"] == [
"feature-b"
]
def test_normalize_operation_ids_uses_each_http_method():
paths = {
"/proxy/{endpoint}": {

View file

@ -5,7 +5,7 @@ import os
import socket
import subprocess
import sys
from datetime import datetime, timezone
from datetime import datetime, timedelta, timezone
from pathlib import Path
from unittest import mock
from unittest.mock import AsyncMock, MagicMock, mock_open, patch
@ -5084,8 +5084,12 @@ async def test_init_and_increment_spend_counter_reseeds_from_db_on_counter_miss(
@pytest.mark.asyncio
async def test_reseed_spend_from_db_user_and_org_prefixes():
"""User and org counters must reseed from their own DB tables, not
fall through to 0.0 like the other counters do today."""
"""User and org counters reseed from their own DB tables.
End-user and tag counters use the already fetched auth objects passed as
fallback_spend, so this reseed helper must not add extra per-request DB
reads for them.
"""
from litellm.proxy.db.spend_counter_reseed import SpendCounterReseed
user_row = MagicMock()
@ -5095,6 +5099,8 @@ async def test_reseed_spend_from_db_user_and_org_prefixes():
fake_prisma = MagicMock()
fake_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row)
fake_prisma.db.litellm_endusertable.find_unique = AsyncMock()
fake_prisma.db.litellm_tagtable.find_unique = AsyncMock()
fake_prisma.db.litellm_organizationtable.find_unique = AsyncMock(
return_value=org_row
)
@ -5104,6 +5110,18 @@ async def test_reseed_spend_from_db_user_and_org_prefixes():
where={"user_id": "alice"}
)
assert (
await SpendCounterReseed.from_db(
fake_prisma,
"spend:end_user:customer-1",
)
is None
)
fake_prisma.db.litellm_endusertable.find_unique.assert_not_awaited()
assert await SpendCounterReseed.from_db(fake_prisma, "spend:tag:paid-tag") is None
fake_prisma.db.litellm_tagtable.find_unique.assert_not_awaited()
assert await SpendCounterReseed.from_db(fake_prisma, "spend:org:acme") == 305.0
fake_prisma.db.litellm_organizationtable.find_unique.assert_awaited_once_with(
where={"organization_id": "acme"}
@ -5133,6 +5151,468 @@ async def test_reseed_spend_from_db_skips_window_variant_keys():
fake_prisma.db.litellm_teamtable.find_unique.assert_not_awaited()
@pytest.mark.asyncio
async def test_window_spend_counter_reseeds_from_spend_logs_on_counter_miss():
from litellm.caching.dual_cache import DualCache
from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter
counter_cache = DualCache()
window_start = datetime.now(timezone.utc) - timedelta(hours=1)
fake_prisma = MagicMock()
fake_prisma.db.litellm_spendlogs.group_by = AsyncMock(
return_value=[{"api_key": "key-window", "_sum": {"spend": 2.25}}]
)
import litellm.proxy.proxy_server as ps
orig_counter, orig_prisma = ps.spend_counter_cache, ps.prisma_client
ps.spend_counter_cache = counter_cache
ps.prisma_client = fake_prisma
try:
await _init_and_increment_window_spend_counter(
counter_key="spend:key:key-window:window:1h",
entity_type="Key",
entity_id="key-window",
window_start=window_start,
increment=0.5,
)
fake_prisma.db.litellm_spendlogs.group_by.assert_awaited_once_with(
by=["api_key"],
where={"api_key": "key-window", "startTime": {"gte": window_start}},
sum={"spend": True},
)
assert counter_cache.in_memory_cache.get_cache(
key="spend:key:key-window:window:1h"
) == pytest.approx(2.75)
finally:
ps.spend_counter_cache = orig_counter
ps.prisma_client = orig_prisma
@pytest.mark.asyncio
async def test_init_spend_counter_redis_clean_miss_skips_stale_in_memory():
from litellm.caching.dual_cache import DualCache
from litellm.proxy.proxy_server import _init_and_increment_spend_counter
counter_cache = DualCache()
counter_key = "spend:team:team-stale-local"
counter_cache.in_memory_cache.set_cache(key=counter_key, value=10.0)
redis_store: dict = {}
async def redis_increment(key, value, **_):
redis_store[key] = (redis_store.get(key) or 0.0) + value
return redis_store[key]
fake_redis = AsyncMock()
fake_redis.async_get_cache = AsyncMock(return_value=None)
fake_redis.async_increment = AsyncMock(side_effect=redis_increment)
counter_cache.redis_cache = fake_redis
db_row = MagicMock()
db_row.spend = 42.0
fake_prisma = MagicMock()
fake_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=db_row)
import litellm.proxy.proxy_server as ps
orig_counter, orig_prisma, orig_user = (
ps.spend_counter_cache,
ps.prisma_client,
ps.user_api_key_cache,
)
ps.spend_counter_cache = counter_cache
ps.prisma_client = fake_prisma
ps.user_api_key_cache = DualCache()
try:
await _init_and_increment_spend_counter(
counter_key=counter_key,
source_cache_key="team_id:team-stale-local",
increment=1.5,
)
fake_prisma.db.litellm_teamtable.find_unique.assert_awaited_once_with(
where={"team_id": "team-stale-local"}
)
assert redis_store[counter_key] == pytest.approx(43.5)
assert counter_cache.in_memory_cache.get_cache(
key=counter_key
) == pytest.approx(43.5)
finally:
ps.spend_counter_cache = orig_counter
ps.prisma_client = orig_prisma
ps.user_api_key_cache = orig_user
@pytest.mark.asyncio
async def test_window_spend_counter_redis_clean_miss_skips_stale_in_memory():
from litellm.caching.dual_cache import DualCache
from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter
counter_cache = DualCache()
counter_key = "spend:key:key-window-stale-local:window:1h"
counter_cache.in_memory_cache.set_cache(key=counter_key, value=100.0)
window_start = datetime.now(timezone.utc) - timedelta(hours=1)
redis_store: dict = {}
async def redis_increment(key, value, **_):
redis_store[key] = (redis_store.get(key) or 0.0) + value
return redis_store[key]
async def redis_set_cache(key, value, **_):
if key in redis_store:
return False
redis_store[key] = value
return True
fake_redis = AsyncMock()
fake_redis.async_get_cache = AsyncMock(return_value=None)
fake_redis.async_set_cache = AsyncMock(side_effect=redis_set_cache)
fake_redis.async_increment = AsyncMock(side_effect=redis_increment)
counter_cache.redis_cache = fake_redis
fake_prisma = MagicMock()
fake_prisma.db.litellm_spendlogs.group_by = AsyncMock(
return_value=[{"api_key": "key-window-stale-local", "_sum": {"spend": 2.25}}]
)
import litellm.proxy.proxy_server as ps
orig_counter, orig_prisma = ps.spend_counter_cache, ps.prisma_client
ps.spend_counter_cache = counter_cache
ps.prisma_client = fake_prisma
try:
await _init_and_increment_window_spend_counter(
counter_key=counter_key,
entity_type="Key",
entity_id="key-window-stale-local",
window_start=window_start,
increment=0.5,
)
fake_prisma.db.litellm_spendlogs.group_by.assert_awaited_once_with(
by=["api_key"],
where={
"api_key": "key-window-stale-local",
"startTime": {"gte": window_start},
},
sum={"spend": True},
)
assert redis_store[counter_key] == pytest.approx(2.75)
assert counter_cache.in_memory_cache.get_cache(
key=counter_key
) == pytest.approx(2.75)
finally:
ps.spend_counter_cache = orig_counter
ps.prisma_client = orig_prisma
@pytest.mark.asyncio
async def test_window_spend_counter_redis_concurrent_seed_does_not_double_seed():
from litellm.caching.dual_cache import DualCache
from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter
counter_cache = DualCache()
counter_key = "spend:key:key-window-concurrent-seed:window:1h"
window_start = datetime.now(timezone.utc) - timedelta(hours=1)
redis_store = {counter_key: 2.75}
redis_reads = 0
async def redis_get_cache(key):
nonlocal redis_reads
redis_reads += 1
if redis_reads <= 2:
return None
return redis_store.get(key)
async def redis_increment(key, value, **_):
redis_store[key] = (redis_store.get(key) or 0.0) + value
return redis_store[key]
fake_redis = AsyncMock()
fake_redis.async_get_cache = AsyncMock(side_effect=redis_get_cache)
fake_redis.async_set_cache = AsyncMock(return_value=False)
fake_redis.async_increment = AsyncMock(side_effect=redis_increment)
counter_cache.redis_cache = fake_redis
fake_prisma = MagicMock()
fake_prisma.db.litellm_spendlogs.group_by = AsyncMock(
return_value=[
{"api_key": "key-window-concurrent-seed", "_sum": {"spend": 2.25}}
]
)
import litellm.proxy.proxy_server as ps
orig_counter, orig_prisma = ps.spend_counter_cache, ps.prisma_client
ps.spend_counter_cache = counter_cache
ps.prisma_client = fake_prisma
try:
await _init_and_increment_window_spend_counter(
counter_key=counter_key,
entity_type="Key",
entity_id="key-window-concurrent-seed",
window_start=window_start,
increment=0.5,
)
fake_redis.async_set_cache.assert_awaited_once_with(
key=counter_key,
value=2.25,
nx=True,
)
assert redis_store[counter_key] == pytest.approx(3.25)
assert counter_cache.in_memory_cache.get_cache(
key=counter_key
) == pytest.approx(3.25)
finally:
ps.spend_counter_cache = orig_counter
ps.prisma_client = orig_prisma
@pytest.mark.asyncio
async def test_window_spend_counter_skips_invalid_window_start():
from litellm.caching.dual_cache import DualCache
from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter
counter_cache = DualCache()
import litellm.proxy.proxy_server as ps
orig_counter = ps.spend_counter_cache
ps.spend_counter_cache = counter_cache
try:
await _init_and_increment_window_spend_counter(
counter_key="spend:key:key-invalid-window:window:not-a-duration",
entity_type="Key",
entity_id="key-invalid-window",
window_start=None,
increment=0.5,
)
assert (
counter_cache.in_memory_cache.get_cache(
key="spend:key:key-invalid-window:window:not-a-duration"
)
is None
)
finally:
ps.spend_counter_cache = orig_counter
@pytest.mark.asyncio
async def test_window_spend_counter_does_not_seed_zero_when_db_unavailable():
from litellm.caching.dual_cache import DualCache
from litellm.proxy.proxy_server import _ensure_window_spend_counter_initialized
counter_cache = DualCache()
counter_key = "spend:key:key-window-db-unavailable:window:1h"
import litellm.proxy.proxy_server as ps
orig_counter, orig_prisma = ps.spend_counter_cache, ps.prisma_client
ps.spend_counter_cache = counter_cache
ps.prisma_client = None
try:
initialized = await _ensure_window_spend_counter_initialized(
counter_key=counter_key,
entity_type="Key",
entity_id="key-window-db-unavailable",
window_start=datetime.now(timezone.utc) - timedelta(hours=1),
)
assert initialized is False
assert counter_cache.in_memory_cache.get_cache(key=counter_key) is None
finally:
ps.spend_counter_cache = orig_counter
ps.prisma_client = orig_prisma
@pytest.mark.asyncio
async def test_increment_spend_counters_finalizes_after_unreserved_increments():
from litellm.caching.dual_cache import DualCache
from litellm.proxy.proxy_server import increment_spend_counters
counter_cache = DualCache()
counter_cache.in_memory_cache.set_cache(
key="spend:key:key-finalize-after-increments",
value=0.5,
)
budget_reservation = {
"reserved_cost": 0.5,
"entries": [
{
"counter_key": "spend:key:key-finalize-after-increments",
"entity_type": "Key",
"entity_id": "key-finalize-after-increments",
"reserved_cost": 0.5,
"applied_adjustment": 0.0,
}
],
"finalized": False,
}
incremented_counters = []
async def assert_reservation_not_finalized_yet(**kwargs):
assert budget_reservation["finalized"] is False
incremented_counters.append(kwargs["counter_key"])
import litellm.proxy.proxy_server as ps
orig_counter, orig_user = ps.spend_counter_cache, ps.user_api_key_cache
ps.spend_counter_cache = counter_cache
ps.user_api_key_cache = DualCache()
try:
with patch(
"litellm.proxy.proxy_server._init_and_increment_spend_counter",
new=AsyncMock(side_effect=assert_reservation_not_finalized_yet),
):
await increment_spend_counters(
token="key-finalize-after-increments",
team_id="team-finalize-after-increments",
user_id=None,
response_cost=0.25,
budget_reservation=budget_reservation,
)
assert incremented_counters == ["spend:team:team-finalize-after-increments"]
assert budget_reservation["finalized"] is True
assert counter_cache.in_memory_cache.get_cache(
key="spend:key:key-finalize-after-increments"
) == pytest.approx(0.25)
finally:
ps.spend_counter_cache = orig_counter
ps.user_api_key_cache = orig_user
@pytest.mark.asyncio
async def test_increment_spend_counters_finalizes_none_cost_reservation():
from litellm.caching.dual_cache import DualCache
from litellm.proxy.proxy_server import increment_spend_counters
counter_cache = DualCache()
counter_cache.in_memory_cache.set_cache(
key="spend:key:key-finalize-none-cost",
value=0.5,
)
budget_reservation = {
"reserved_cost": 0.5,
"entries": [
{
"counter_key": "spend:key:key-finalize-none-cost",
"entity_type": "Key",
"entity_id": "key-finalize-none-cost",
"reserved_cost": 0.5,
"applied_adjustment": 0.0,
}
],
"finalized": False,
}
import litellm.proxy.proxy_server as ps
orig_counter = ps.spend_counter_cache
ps.spend_counter_cache = counter_cache
try:
await increment_spend_counters(
token="key-finalize-none-cost",
team_id=None,
user_id=None,
response_cost=None,
budget_reservation=budget_reservation,
)
assert budget_reservation["finalized"] is True
assert counter_cache.in_memory_cache.get_cache(
key="spend:key:key-finalize-none-cost"
) == pytest.approx(0.0)
finally:
ps.spend_counter_cache = orig_counter
@pytest.mark.asyncio
async def test_increment_spend_counters_invalidates_bad_reserved_counter_without_failing():
from litellm.caching.dual_cache import DualCache
from litellm.proxy.proxy_server import increment_spend_counters
counter_cache = DualCache()
budget_reservation = {
"reserved_cost": 0.5,
"entries": [
{
"counter_key": "spend:key:key-bad-reserved-counter",
"entity_type": "Key",
"entity_id": "key-bad-reserved-counter",
"reserved_cost": 0.5,
"applied_adjustment": 0.0,
}
],
"finalized": False,
}
import litellm.proxy.proxy_server as ps
orig_counter = ps.spend_counter_cache
ps.spend_counter_cache = counter_cache
try:
with patch(
"litellm.proxy.proxy_server.verbose_proxy_logger.warning"
) as mock_warning:
await increment_spend_counters(
token="key-bad-reserved-counter",
team_id=None,
user_id=None,
response_cost=0.25,
budget_reservation=budget_reservation,
)
mock_warning.assert_called_once()
assert budget_reservation["finalized"] is True
assert (
counter_cache.in_memory_cache.get_cache(
key="spend:key:key-bad-reserved-counter"
)
is None
)
finally:
ps.spend_counter_cache = orig_counter
@pytest.mark.asyncio
async def test_increment_spend_counter_invalidates_stale_cache_on_redis_failure():
from litellm.caching.dual_cache import DualCache
from litellm.proxy.proxy_server import _increment_spend_counter_cache
counter_cache = DualCache()
counter_cache.in_memory_cache.set_cache(key="spend:team:redis-fail", value=4.0)
fake_redis = AsyncMock()
fake_redis.async_increment = AsyncMock(side_effect=RuntimeError("redis down"))
fake_redis.async_delete_cache = AsyncMock()
counter_cache.redis_cache = fake_redis
import litellm.proxy.proxy_server as ps
orig_counter = ps.spend_counter_cache
ps.spend_counter_cache = counter_cache
try:
with pytest.raises(RuntimeError):
await _increment_spend_counter_cache(
counter_key="spend:team:redis-fail",
increment=0.5,
)
assert (
counter_cache.in_memory_cache.get_cache(key="spend:team:redis-fail") is None
)
fake_redis.async_delete_cache.assert_awaited_once_with(
key="spend:team:redis-fail"
)
finally:
ps.spend_counter_cache = orig_counter
@pytest.mark.asyncio
async def test_get_current_spend_reseeds_from_db_when_counter_missing():
"""
@ -5181,6 +5661,9 @@ async def test_get_current_spend_reseeds_from_db_when_counter_missing():
assert ("spend:team_member:user-1:team-1", 362.0) in [
(w["key"], w["value"]) for w in recorded_warms
]
assert counter_cache.in_memory_cache.get_cache(
key="spend:team_member:user-1:team-1"
) == pytest.approx(362.0)
finally:
ps.spend_counter_cache = orig_counter
ps.prisma_client = orig_prisma

View file

@ -156,8 +156,11 @@ async def test_update_request_data_with_litellm_managed_vector_store_registry():
vector_store_id="test_store_id"
)
# Test with no vector store registry
with patch.object(litellm, "vector_store_registry", None):
# Test with no vector store registry or DB fallback
with (
patch.object(litellm, "vector_store_registry", None),
patch("litellm.proxy.proxy_server.prisma_client", None),
):
original_data = {"existing_key": "existing_value"}
result = await _update_request_data_with_litellm_managed_vector_store_registry(
data=original_data, vector_store_id=vector_store_id

View file

@ -0,0 +1,540 @@
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException, Request, Response
import litellm
from litellm.proxy._types import LiteLLM_ManagedVectorStoresTable, UserAPIKeyAuth
def _mock_request() -> MagicMock:
request = MagicMock(spec=Request)
request.headers = {}
request.method = "POST"
request.query_params = {}
request.url.path = "/v1/vector_stores/vs_path/search"
return request
@pytest.mark.asyncio
async def test_vector_store_search_forces_path_id_over_body_id():
from litellm.proxy.vector_store_endpoints.endpoints import vector_store_search
captured_data = {}
async def fake_base_process(self, **kwargs):
captured_data.update(self.data)
return {"ok": True}
request = _mock_request()
with (
patch(
"litellm.proxy.proxy_server._read_request_body",
new=AsyncMock(
return_value={
"vector_store_id": "vs_body_victim",
"query": "test",
}
),
),
patch.object(litellm, "vector_store_registry", None),
patch("litellm.proxy.proxy_server.prisma_client", None),
patch(
"litellm.proxy.vector_store_endpoints.endpoints.ProxyBaseLLMRequestProcessing.base_process_llm_request",
new=fake_base_process,
),
):
response = await vector_store_search(
request=request,
vector_store_id="vs_path_allowed",
fastapi_response=Response(),
user_api_key_dict=UserAPIKeyAuth(team_id="team-a"),
)
assert response == {"ok": True}
assert captured_data["vector_store_id"] == "vs_path_allowed"
@pytest.mark.asyncio
async def test_vector_store_file_create_forces_path_id_over_body_id():
from litellm.proxy.vector_store_files_endpoints.endpoints import (
vector_store_file_create,
)
captured_data = {}
async def fake_base_process(self, **kwargs):
captured_data.update(self.data)
return {"ok": True}
mock_registry = MagicMock()
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = {
"vector_store_id": "vs_path_allowed",
"custom_llm_provider": "openai",
"team_id": "team-a",
}
request = _mock_request()
with (
patch(
"litellm.proxy.proxy_server._read_request_body",
new=AsyncMock(
return_value={
"vector_store_id": "vs_body_victim",
"file_id": "file_123",
}
),
),
patch.object(litellm, "vector_store_registry", mock_registry),
patch(
"litellm.proxy.vector_store_files_endpoints.endpoints.ProxyBaseLLMRequestProcessing.base_process_llm_request",
new=fake_base_process,
),
):
response = await vector_store_file_create(
vector_store_id="vs_path_allowed",
request=request,
fastapi_response=Response(),
user_api_key_dict=UserAPIKeyAuth(team_id="team-a"),
)
assert response == {"ok": True}
assert captured_data["vector_store_id"] == "vs_path_allowed"
assert captured_data["custom_llm_provider"] == "openai"
mock_registry.get_litellm_managed_vector_store_from_registry.assert_called_once_with(
vector_store_id="vs_path_allowed"
)
@pytest.mark.asyncio
async def test_vector_store_file_create_denies_other_team_path_store():
from litellm.proxy.vector_store_files_endpoints.endpoints import (
vector_store_file_create,
)
mock_registry = MagicMock()
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = {
"vector_store_id": "vs_other_team",
"custom_llm_provider": "openai",
"team_id": "team-b",
}
request = _mock_request()
with (
patch(
"litellm.proxy.proxy_server._read_request_body",
new=AsyncMock(return_value={"file_id": "file_123"}),
),
patch.object(litellm, "vector_store_registry", mock_registry),
patch(
"litellm.proxy.vector_store_files_endpoints.endpoints.ProxyBaseLLMRequestProcessing.base_process_llm_request",
new=AsyncMock(),
) as mock_base_process,
):
with pytest.raises(HTTPException) as exc_info:
await vector_store_file_create(
vector_store_id="vs_other_team",
request=request,
fastapi_response=Response(),
user_api_key_dict=UserAPIKeyAuth(team_id="team-a"),
)
assert exc_info.value.status_code == 403
mock_base_process.assert_not_called()
@pytest.mark.asyncio
async def test_rag_query_denies_nested_other_team_vector_store():
from litellm.proxy.rag_endpoints.endpoints import rag_query
mock_registry = MagicMock()
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = {
"vector_store_id": "vs_other_team",
"custom_llm_provider": "openai",
"team_id": "team-b",
}
request = _mock_request()
with (
patch(
"litellm.proxy.rag_endpoints.endpoints._read_request_body",
new=AsyncMock(
return_value={
"model": "gpt-4o-mini",
"messages": [{"role": "user", "content": "hello"}],
"retrieval_config": {"vector_store_id": "vs_other_team"},
}
),
),
patch.object(litellm, "vector_store_registry", mock_registry),
patch(
"litellm.proxy.rag_endpoints.endpoints.litellm.aquery",
new=AsyncMock(),
) as mock_aquery,
):
with pytest.raises(HTTPException) as exc_info:
await rag_query(
request=request,
fastapi_response=Response(),
user_api_key_dict=UserAPIKeyAuth(team_id="team-a"),
)
assert exc_info.value.status_code == 403
mock_aquery.assert_not_called()
@pytest.mark.asyncio
async def test_rag_ingest_denies_nested_other_team_vector_store():
from litellm.proxy.rag_endpoints.endpoints import rag_ingest
mock_registry = MagicMock()
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = {
"vector_store_id": "vs_other_team",
"custom_llm_provider": "openai",
"team_id": "team-b",
}
request = _mock_request()
with (
patch(
"litellm.proxy.rag_endpoints.endpoints.parse_rag_ingest_request",
new=AsyncMock(
return_value=(
{
"vector_store": {
"custom_llm_provider": "openai",
"vector_store_id": "vs_other_team",
}
},
None,
"https://example.com/file.txt",
None,
)
),
),
patch.object(litellm, "vector_store_registry", mock_registry),
patch(
"litellm.proxy.rag_endpoints.endpoints.litellm.aingest",
new=AsyncMock(),
) as mock_aingest,
):
with pytest.raises(HTTPException) as exc_info:
await rag_ingest(
request=request,
fastapi_response=Response(),
user_api_key_dict=UserAPIKeyAuth(team_id="team-a"),
)
assert exc_info.value.status_code == 403
mock_aingest.assert_not_called()
def test_rag_payload_scan_rejects_excessive_nesting():
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
from litellm.proxy.rag_endpoints.endpoints import (
_collect_vector_store_ids_from_payload,
)
payload = {}
current = payload
for _ in range(DEFAULT_MAX_RECURSE_DEPTH + 1):
current["nested"] = {}
current = current["nested"]
current["vector_store_id"] = "vs_too_deep"
with pytest.raises(HTTPException) as exc_info:
_collect_vector_store_ids_from_payload(payload)
assert exc_info.value.status_code == 400
def test_rag_payload_scan_accepts_vector_store_id_at_depth_limit():
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
from litellm.proxy.rag_endpoints.endpoints import (
_collect_vector_store_ids_from_payload,
)
payload = {}
current = payload
for _ in range(DEFAULT_MAX_RECURSE_DEPTH):
current["nested"] = {}
current = current["nested"]
current["vector_store_id"] = "vs_at_limit"
assert _collect_vector_store_ids_from_payload(payload) == {"vs_at_limit"}
def test_rag_payload_scan_ignores_primitive_list_beyond_depth_limit():
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
from litellm.proxy.rag_endpoints.endpoints import (
_collect_vector_store_ids_from_payload,
)
payload = {}
current = payload
for _ in range(DEFAULT_MAX_RECURSE_DEPTH):
current["nested"] = {}
current = current["nested"]
current["labels"] = ["alpha", "beta"]
assert _collect_vector_store_ids_from_payload(payload) == set()
@pytest.mark.asyncio
async def test_responses_file_search_denies_other_team_vector_store():
from litellm.proxy.common_request_processing import (
_authorize_response_file_search_vector_stores,
)
mock_registry = MagicMock()
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = {
"vector_store_id": "vs_other_team",
"custom_llm_provider": "openai",
"team_id": "team-b",
}
with patch.object(litellm, "vector_store_registry", mock_registry):
with pytest.raises(HTTPException) as exc_info:
await _authorize_response_file_search_vector_stores(
data={
"tools": [
{
"type": "file_search",
"vector_store_ids": ["vs_other_team"],
}
]
},
user_api_key_dict=UserAPIKeyAuth(team_id="team-a"),
)
assert exc_info.value.status_code == 403
@pytest.mark.asyncio
async def test_vertex_discovery_denies_other_team_vector_store_credentials():
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
_base_vertex_proxy_route,
)
request = _mock_request()
request.method = "GET"
vector_store_credentials = {
"vector_store_id": "vs_other_team",
"custom_llm_provider": "vertex_ai",
"team_id": "team-b",
}
with patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.user_api_key_auth",
new=AsyncMock(return_value=UserAPIKeyAuth(team_id="team-a")),
):
with pytest.raises(HTTPException) as exc_info:
await _base_vertex_proxy_route(
endpoint="projects/p/locations/us-central1/dataStores/vs_other_team",
request=request,
fastapi_response=Response(),
get_vertex_pass_through_handler=MagicMock(),
router_credentials=vector_store_credentials,
)
assert exc_info.value.status_code == 403
@pytest.mark.asyncio
async def test_get_managed_vector_store_uses_shared_cache_helper_for_db_fallback():
from litellm.proxy.vector_store_endpoints.utils import (
get_litellm_managed_vector_store,
)
mock_registry = MagicMock()
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = None
cache_helper = AsyncMock(
return_value=[
LiteLLM_ManagedVectorStoresTable(
vector_store_id="vs_cached",
custom_llm_provider="openai",
vector_store_name=None,
vector_store_description=None,
vector_store_metadata=None,
created_at=None,
updated_at=None,
litellm_credential_name=None,
litellm_params={"api_base": "https://example.com"},
team_id="team-a",
user_id=None,
)
]
)
with (
patch.object(litellm, "vector_store_registry", mock_registry),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
patch(
"litellm.proxy.auth.auth_checks.get_managed_vector_store_rows_by_uuids",
new=cache_helper,
),
):
vector_store = await get_litellm_managed_vector_store(
vector_store_id="vs_cached"
)
assert vector_store is not None
assert vector_store["vector_store_id"] == "vs_cached"
assert vector_store["team_id"] == "team-a"
cache_helper.assert_awaited_once()
@pytest.mark.asyncio
async def test_get_managed_vector_store_fails_closed_on_lookup_error():
from litellm.proxy.vector_store_endpoints.utils import (
get_litellm_managed_vector_store,
)
mock_registry = MagicMock()
mock_registry.get_litellm_managed_vector_store_from_registry.side_effect = (
RuntimeError("registry unavailable")
)
with patch.object(litellm, "vector_store_registry", mock_registry):
with pytest.raises(HTTPException) as exc_info:
await get_litellm_managed_vector_store(vector_store_id="vs_registry_only")
assert exc_info.value.status_code == 500
@pytest.mark.asyncio
async def test_vertex_discovery_allows_unregistered_provider_native_datastore_id():
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
vertex_discovery_proxy_route,
)
request = _mock_request()
request.method = "GET"
with (
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_litellm_managed_vector_store",
new=AsyncMock(return_value=None),
) as mock_lookup,
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._base_vertex_proxy_route",
new=AsyncMock(return_value={"ok": True}),
) as mock_base_route,
):
response = await vertex_discovery_proxy_route(
endpoint="projects/p/locations/us-central1/dataStores/vs_unknown",
request=request,
fastapi_response=Response(),
)
assert response == {"ok": True}
mock_lookup.assert_awaited_once_with(vector_store_id="vs_unknown")
assert mock_base_route.call_args.kwargs["router_credentials"] is None
@pytest.mark.asyncio
async def test_milvus_passthrough_denies_other_team_vector_store_index():
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
milvus_proxy_route,
)
request = _mock_request()
request.url.path = "/milvus/v2/vectordb/entities/search"
index_object = MagicMock()
index_object.litellm_params.vector_store_name = "tenant-b-store"
index_object.litellm_params.vector_store_index = "tenant_b_collection"
mock_index_registry = MagicMock()
mock_index_registry.is_vector_store_index.return_value = True
mock_index_registry.get_vector_store_index_by_name.return_value = index_object
mock_vector_registry = MagicMock()
mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = {
"vector_store_id": "vs_other_team",
"custom_llm_provider": "milvus",
"team_id": "team-b",
"litellm_params": {"api_base": "https://milvus.example.com"},
}
with (
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_vector_stores_config",
return_value=MagicMock(),
),
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_request_body",
new=AsyncMock(return_value={"collectionName": "managed_index"}),
),
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint",
return_value=True,
),
patch.object(litellm, "vector_store_index_registry", mock_index_registry),
patch.object(litellm, "vector_store_registry", mock_vector_registry),
):
with pytest.raises(HTTPException) as exc_info:
await milvus_proxy_route(
endpoint="v2/vectordb/entities/search",
request=request,
fastapi_response=Response(),
user_api_key_dict=UserAPIKeyAuth(team_id="team-a"),
)
assert exc_info.value.status_code == 403
@pytest.mark.asyncio
async def test_azure_passthrough_denies_other_team_vector_store_index():
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
azure_proxy_route,
)
request = _mock_request()
request.url.path = "/azure/indexes/managed_index/docs/search"
index_object = MagicMock()
index_object.litellm_params.vector_store_name = "tenant-b-store"
mock_index_registry = MagicMock()
mock_index_registry.is_vector_store_index.side_effect = (
lambda vector_store_index_name: vector_store_index_name == "managed_index"
)
mock_index_registry.get_vector_store_index_by_name.return_value = index_object
mock_vector_registry = MagicMock()
mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = {
"vector_store_id": "vs_other_team",
"custom_llm_provider": "azure_ai",
"team_id": "team-b",
"litellm_params": {"api_base": "https://azure.example.com"},
}
with (
patch("litellm.proxy.proxy_server.llm_router", MagicMock()),
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_passthrough_request_using_router_model",
return_value=False,
),
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_vector_stores_config",
return_value=MagicMock(),
),
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint",
return_value=True,
),
patch.object(litellm, "vector_store_index_registry", mock_index_registry),
patch.object(litellm, "vector_store_registry", mock_vector_registry),
):
with pytest.raises(HTTPException) as exc_info:
await azure_proxy_route(
endpoint="indexes/managed_index/docs/search",
request=request,
fastapi_response=Response(),
user_api_key_dict=UserAPIKeyAuth(team_id="team-a"),
)
assert exc_info.value.status_code == 403