mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge pull request #27024 from BerriAI/litellm_yj_may1
[Infra] Merge dev branch
This commit is contained in:
commit
57dd3891fb
39 changed files with 6991 additions and 396 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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", {})},
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}",
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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] = {}
|
||||
|
|
|
|||
1029
litellm/proxy/spend_tracking/budget_reservation.py
Normal file
1029
litellm/proxy/spend_tracking/budget_reservation.py
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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"
|
||||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
208
tests/test_litellm/proxy/hooks/test_max_budget_limiter.py
Normal file
208
tests/test_litellm/proxy/hooks/test_max_budget_limiter.py
Normal 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()
|
||||
|
|
@ -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()
|
||||
|
||||
|
||||
|
|
|
|||
1495
tests/test_litellm/proxy/test_budget_reservation.py
Normal file
1495
tests/test_litellm/proxy/test_budget_reservation.py
Normal file
File diff suppressed because it is too large
Load diff
102
tests/test_litellm/proxy/test_langfuse_passthrough_security.py
Normal file
102
tests/test_litellm/proxy/test_langfuse_passthrough_security.py
Normal 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"
|
||||
|
|
@ -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}": {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
Loading…
Add table
Reference in a new issue