mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge branch 'litellm_yj_may1' into codex/budget-race-enforcement
This commit is contained in:
commit
c2cea58567
93 changed files with 4461 additions and 579 deletions
|
|
@ -1,75 +0,0 @@
|
|||
name: Check Lazy OpenAPI Snapshot
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- "litellm_**"
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
checks: write
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
verify:
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 10
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Set up uv
|
||||
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
|
||||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
- name: Cache uv dependencies
|
||||
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
|
||||
with:
|
||||
path: |
|
||||
~/.cache/uv
|
||||
.venv
|
||||
key: ${{ runner.os }}-uv-${{ hashFiles('uv.lock') }}
|
||||
restore-keys: |
|
||||
${{ runner.os }}-uv-
|
||||
|
||||
- name: Install dependencies
|
||||
run: uv sync --frozen --all-groups --all-extras
|
||||
|
||||
- name: Regenerate snapshot to /tmp
|
||||
id: regen
|
||||
run: |
|
||||
cp litellm/proxy/_lazy_openapi_snapshot.json /tmp/snapshot.committed.json
|
||||
uv run --no-sync python -m litellm.proxy._lazy_openapi_snapshot
|
||||
mv litellm/proxy/_lazy_openapi_snapshot.json /tmp/snapshot.fresh.json
|
||||
mv /tmp/snapshot.committed.json litellm/proxy/_lazy_openapi_snapshot.json
|
||||
|
||||
- name: Compare
|
||||
id: diff
|
||||
continue-on-error: true
|
||||
run: |
|
||||
diff -q /tmp/snapshot.fresh.json litellm/proxy/_lazy_openapi_snapshot.json
|
||||
|
||||
- name: Mark neutral if drift
|
||||
if: steps.diff.outcome == 'failure'
|
||||
uses: LouisBrunner/checks-action@6b626ffbad7cc56fd58627f774b9067e6118af23 # v2.0.0
|
||||
with:
|
||||
token: ${{ secrets.GITHUB_TOKEN }}
|
||||
name: lazy-openapi-snapshot
|
||||
conclusion: neutral
|
||||
output: |
|
||||
{
|
||||
"title": "Lazy openapi snapshot is stale",
|
||||
"summary": "Run `python -m litellm.proxy._lazy_openapi_snapshot` and commit the regenerated `litellm/proxy/_lazy_openapi_snapshot.json`. Not blocking — the snapshot will regenerate at release if not committed."
|
||||
}
|
||||
|
|
@ -68,7 +68,7 @@ Managing LLM calls across providers gets complicated fast — different SDKs, au
|
|||
<td><img height="60" alt="Stripe" src="https://github.com/user-attachments/assets/f7296d4f-9fbd-460d-9d05-e4df31697c4b" /></td>
|
||||
<td><img height="60" alt="image" src="https://github.com/user-attachments/assets/436fca71-988b-40bb-b5fe-8450c80fdbd0" /></td>
|
||||
<td><img height="60" alt="Google ADK" src="https://github.com/user-attachments/assets/caf270a2-5aee-45c4-8222-41a2070c4f19" /></td>
|
||||
<td><img height="60" alt="Greptile" src="https://github.com/user-attachments/assets/0be4bd8a-7cfa-48d3-9090-f415fe948280" /></td>
|
||||
<td><img height="60" alt="Greptile" src="https://github.com/user-attachments/assets/3db0ae72-0843-4005-a56d-bba1dde2193d" /></td>
|
||||
<td><img height="60" alt="OpenHands" src="https://github.com/user-attachments/assets/a6150c4c-149e-4cae-888b-8b92be6e003f" /></td>
|
||||
<td><h2>Netflix</h2></td>
|
||||
<td><img height="60" alt="OpenAI Agents SDK" src="https://github.com/user-attachments/assets/c02f7be0-8c2e-4d27-aea7-7c024bfaebc0" /></td>
|
||||
|
|
|
|||
|
|
@ -857,10 +857,16 @@ async def project_info(
|
|||
where={"team_id": project.team_id}
|
||||
)
|
||||
if team:
|
||||
is_team_member = (
|
||||
user_api_key_dict.user_id in team.admins
|
||||
or user_api_key_dict.user_id in team.members
|
||||
)
|
||||
caller_user_id = user_api_key_dict.user_id
|
||||
for m in team.members_with_roles or []:
|
||||
m_user_id = (
|
||||
m.get("user_id")
|
||||
if isinstance(m, dict)
|
||||
else getattr(m, "user_id", None)
|
||||
)
|
||||
if m_user_id == caller_user_id:
|
||||
is_team_member = True
|
||||
break
|
||||
|
||||
if not (is_admin or is_team_member):
|
||||
raise HTTPException(
|
||||
|
|
@ -911,20 +917,20 @@ async def list_projects(
|
|||
include={"litellm_budget_table": True, "object_permission": True}
|
||||
)
|
||||
else:
|
||||
# Get projects for teams the user belongs to
|
||||
user_teams = await prisma_client.db.litellm_teamtable.find_many(
|
||||
where={
|
||||
"OR": [
|
||||
{"members": {"has": user_api_key_dict.user_id}},
|
||||
{"admins": {"has": user_api_key_dict.user_id}},
|
||||
]
|
||||
}
|
||||
# Look up the user's team memberships via the reverse-index on
|
||||
# LiteLLM_UserTable.teams (maintained by team_member_add alongside
|
||||
# members_with_roles). This avoids a full scan of all team rows.
|
||||
user_record = await prisma_client.db.litellm_usertable.find_unique(
|
||||
where={"user_id": user_api_key_dict.user_id},
|
||||
)
|
||||
user_team_ids = (
|
||||
user_record.teams
|
||||
if user_record is not None and user_record.teams
|
||||
else []
|
||||
)
|
||||
|
||||
team_ids = [team.team_id for team in user_teams]
|
||||
|
||||
projects = await prisma_client.db.litellm_projecttable.find_many(
|
||||
where={"team_id": {"in": team_ids}},
|
||||
where={"team_id": {"in": user_team_ids}},
|
||||
include={"litellm_budget_table": True, "object_permission": True},
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -2,11 +2,23 @@
|
|||
Arize Phoenix API client for fetching prompt versions from Arize Phoenix.
|
||||
"""
|
||||
|
||||
import urllib.parse
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
|
||||
|
||||
def _sanitize_id(identifier: str) -> str:
|
||||
"""Reject path traversal characters and URL-encode the identifier."""
|
||||
if any(c in identifier for c in ("/", "\\", "#", "?")):
|
||||
raise ValueError(
|
||||
f"Invalid identifier {identifier!r}: contains disallowed characters"
|
||||
)
|
||||
if ".." in identifier:
|
||||
raise ValueError(f"Invalid identifier {identifier!r}: path traversal detected")
|
||||
return urllib.parse.quote(identifier, safe="")
|
||||
|
||||
|
||||
class ArizePhoenixClient:
|
||||
"""
|
||||
Client for interacting with Arize Phoenix API to fetch prompt versions.
|
||||
|
|
@ -53,7 +65,8 @@ class ArizePhoenixClient:
|
|||
Returns:
|
||||
Dictionary containing prompt version data, or None if not found
|
||||
"""
|
||||
url = f"{self.api_base}/v1/prompt_versions/{prompt_version_id}"
|
||||
safe_id = _sanitize_id(prompt_version_id)
|
||||
url = f"{self.api_base}/v1/prompt_versions/{safe_id}"
|
||||
|
||||
try:
|
||||
# Use the underlying httpx client directly to avoid query param extraction
|
||||
|
|
|
|||
|
|
@ -3,11 +3,27 @@ BitBucket API client for fetching .prompt files from BitBucket repositories.
|
|||
"""
|
||||
|
||||
import base64
|
||||
import urllib.parse
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
|
||||
|
||||
def _sanitize_file_path(file_path: str) -> str:
|
||||
"""Reject path traversal and URL-encode each path segment."""
|
||||
if "#" in file_path or "?" in file_path:
|
||||
raise ValueError(
|
||||
f"Invalid file path {file_path!r}: contains URL special characters"
|
||||
)
|
||||
parts = file_path.split("/")
|
||||
for part in parts:
|
||||
if part == "..":
|
||||
raise ValueError(
|
||||
f"Invalid file path {file_path!r}: path traversal detected"
|
||||
)
|
||||
return "/".join(urllib.parse.quote(part, safe="") for part in parts)
|
||||
|
||||
|
||||
class BitBucketClient:
|
||||
"""
|
||||
Client for interacting with BitBucket API to fetch .prompt files.
|
||||
|
|
@ -72,7 +88,8 @@ class BitBucketClient:
|
|||
Returns:
|
||||
File content as string, or None if file not found
|
||||
"""
|
||||
url = f"{self.base_url}/repositories/{self.workspace}/{self.repository}/src/{self.branch}/{file_path}"
|
||||
safe_path = _sanitize_file_path(file_path)
|
||||
url = f"{self.base_url}/repositories/{self.workspace}/{self.repository}/src/{self.branch}/{safe_path}"
|
||||
|
||||
try:
|
||||
response = self.http_handler.get(url, headers=self.headers)
|
||||
|
|
@ -119,7 +136,8 @@ class BitBucketClient:
|
|||
Returns:
|
||||
List of file paths
|
||||
"""
|
||||
url = f"{self.base_url}/repositories/{self.workspace}/{self.repository}/src/{self.branch}/{directory_path}"
|
||||
safe_dir = _sanitize_file_path(directory_path) if directory_path else ""
|
||||
url = f"{self.base_url}/repositories/{self.workspace}/{self.repository}/src/{self.branch}/{safe_dir}"
|
||||
|
||||
try:
|
||||
response = self.http_handler.get(url, headers=self.headers)
|
||||
|
|
@ -211,7 +229,8 @@ class BitBucketClient:
|
|||
Returns:
|
||||
Dictionary containing file metadata, or None if file not found
|
||||
"""
|
||||
url = f"{self.base_url}/repositories/{self.workspace}/{self.repository}/src/{self.branch}/{file_path}"
|
||||
safe_path = _sanitize_file_path(file_path)
|
||||
url = f"{self.base_url}/repositories/{self.workspace}/{self.repository}/src/{self.branch}/{safe_path}"
|
||||
|
||||
try:
|
||||
# Use GET with Range header to get just the headers (HEAD equivalent)
|
||||
|
|
|
|||
|
|
@ -31,15 +31,23 @@ def load_cli_token() -> Optional[dict]:
|
|||
return None
|
||||
|
||||
|
||||
def get_litellm_gateway_api_key() -> Optional[str]:
|
||||
def get_litellm_gateway_api_key(
|
||||
expected_base_url: Optional[str] = None,
|
||||
) -> Optional[str]:
|
||||
"""
|
||||
Get the stored CLI API key for use with LiteLLM SDK.
|
||||
|
||||
This function reads the token file created by `litellm-proxy login`
|
||||
and returns the API key for use in Python scripts.
|
||||
|
||||
Args:
|
||||
expected_base_url: When provided, the key is only returned if it was
|
||||
originally issued for this URL. Pass the target server URL to
|
||||
prevent credential leakage when the client is pointed at a
|
||||
different (possibly malicious) server.
|
||||
|
||||
Returns:
|
||||
str: The API key if found, None otherwise
|
||||
str: The API key if found (and origin matches), None otherwise
|
||||
|
||||
Example:
|
||||
>>> import litellm
|
||||
|
|
@ -53,6 +61,10 @@ def get_litellm_gateway_api_key() -> Optional[str]:
|
|||
>>> )
|
||||
"""
|
||||
token_data = load_cli_token()
|
||||
if token_data and "key" in token_data:
|
||||
return token_data["key"]
|
||||
return None
|
||||
if not token_data or "key" not in token_data:
|
||||
return None
|
||||
if expected_base_url is not None:
|
||||
stored_url = token_data.get("base_url")
|
||||
if stored_url != expected_base_url.rstrip("/"):
|
||||
return None
|
||||
return token_data["key"]
|
||||
|
|
|
|||
|
|
@ -2244,7 +2244,7 @@ class CustomStreamWrapper:
|
|||
asyncio.create_task(
|
||||
self.logging_obj.async_failure_handler(e, traceback_exception)
|
||||
)
|
||||
raise e
|
||||
self._handle_stream_fallback_error(e)
|
||||
except Exception as e:
|
||||
traceback_exception = traceback.format_exc()
|
||||
if self.logging_obj is not None:
|
||||
|
|
|
|||
|
|
@ -199,6 +199,47 @@ def validate_url(url: str) -> Tuple[str, str]:
|
|||
return rewritten, host_header
|
||||
|
||||
|
||||
def assert_same_origin(candidate_url: str, expected_url: str) -> None:
|
||||
"""Verify ``candidate_url`` shares scheme, host, and port with ``expected_url``.
|
||||
|
||||
Use when an upstream API returns a URL meant for follow-up requests
|
||||
(e.g. an async-job polling URL that will be hit with the operator's
|
||||
API key in the headers). The upstream is trusted because the operator
|
||||
configured ``api_base``, but the URL it hands back must actually point
|
||||
back at the same origin or we'd be blindly forwarding credentials
|
||||
wherever the upstream told us to.
|
||||
|
||||
Hostnames are compared case-insensitively. Default ports are made
|
||||
explicit (HTTP→80, HTTPS→443) so ``https://api.example.com:443/...``
|
||||
and ``https://api.example.com/...`` are treated as the same origin.
|
||||
|
||||
Error messages identify *which* component mismatched but never echo
|
||||
the operator's ``expected`` host or the candidate's hostname back to
|
||||
the caller — in the SSRF threat model the caller is the attacker,
|
||||
and reflecting host info would be a secondary leak of operator
|
||||
infrastructure details.
|
||||
"""
|
||||
candidate = urlparse(candidate_url)
|
||||
expected = urlparse(expected_url)
|
||||
|
||||
if candidate.scheme not in _ALLOWED_SCHEMES:
|
||||
raise SSRFError("URL scheme is not allowed")
|
||||
|
||||
if candidate.scheme != expected.scheme:
|
||||
raise SSRFError("Origin mismatch on scheme")
|
||||
|
||||
candidate_host = _normalize_host(candidate.hostname or "")
|
||||
expected_host = _normalize_host(expected.hostname or "")
|
||||
if not candidate_host or candidate_host != expected_host:
|
||||
raise SSRFError("Origin mismatch on host")
|
||||
|
||||
default_port = 443 if candidate.scheme == "https" else 80
|
||||
candidate_port = candidate.port if candidate.port is not None else default_port
|
||||
expected_port = expected.port if expected.port is not None else default_port
|
||||
if candidate_port != expected_port:
|
||||
raise SSRFError("Origin mismatch on port")
|
||||
|
||||
|
||||
_MAX_REDIRECTS = 10
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ import litellm
|
|||
from litellm.constants import AZURE_OPERATION_POLLING_TIMEOUT, DEFAULT_MAX_RETRIES
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.logging_utils import track_llm_api_timing
|
||||
from litellm.litellm_core_utils.url_utils import SSRFError, assert_same_origin
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
AsyncHTTPHandler,
|
||||
HTTPHandler,
|
||||
|
|
@ -898,6 +899,17 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
operation_location_url = response.headers["operation-location"]
|
||||
else:
|
||||
raise AzureOpenAIError(status_code=500, message=response.text)
|
||||
# Reject polling URLs that don't share an origin with ``api_base``.
|
||||
# Without this an upstream-controlled or attacker-controlled
|
||||
# value would receive the operator's Azure API key in the
|
||||
# request headers below. VERIA-51.
|
||||
try:
|
||||
assert_same_origin(operation_location_url, api_base)
|
||||
except SSRFError as ssrf_err:
|
||||
raise AzureOpenAIError(
|
||||
status_code=502,
|
||||
message=f"Rejected polling URL: {ssrf_err}",
|
||||
)
|
||||
response = await async_handler.get(
|
||||
url=operation_location_url,
|
||||
headers=headers,
|
||||
|
|
@ -908,8 +920,13 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
timeout_secs: int = AZURE_OPERATION_POLLING_TIMEOUT
|
||||
start_time = time.time()
|
||||
if "status" not in response.json():
|
||||
raise Exception(
|
||||
"Expected 'status' in response. Got={}".format(response.json())
|
||||
# Don't reflect the raw response body — when the polling
|
||||
# URL points at an internal JSON API (cloud metadata
|
||||
# service etc.) reflecting it here turns Blind SSRF into
|
||||
# Full-Read SSRF. VERIA-51.
|
||||
raise AzureOpenAIError(
|
||||
status_code=502,
|
||||
message="Polling response missing 'status' field",
|
||||
)
|
||||
while response.json()["status"] not in ["succeeded", "failed"]:
|
||||
if time.time() - start_time > timeout_secs:
|
||||
|
|
@ -1009,6 +1026,13 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
operation_location_url = response.headers["operation-location"]
|
||||
else:
|
||||
raise AzureOpenAIError(status_code=500, message=response.text)
|
||||
try:
|
||||
assert_same_origin(operation_location_url, api_base)
|
||||
except SSRFError as ssrf_err:
|
||||
raise AzureOpenAIError(
|
||||
status_code=502,
|
||||
message=f"Rejected polling URL: {ssrf_err}",
|
||||
)
|
||||
response = sync_handler.get(
|
||||
url=operation_location_url,
|
||||
headers=headers,
|
||||
|
|
@ -1019,8 +1043,9 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
timeout_secs: int = AZURE_OPERATION_POLLING_TIMEOUT
|
||||
start_time = time.time()
|
||||
if "status" not in response.json():
|
||||
raise Exception(
|
||||
"Expected 'status' in response. Got={}".format(response.json())
|
||||
raise AzureOpenAIError(
|
||||
status_code=502,
|
||||
message="Polling response missing 'status' field",
|
||||
)
|
||||
while response.json()["status"] not in ["succeeded", "failed"]:
|
||||
if time.time() - start_time > timeout_secs:
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ from urllib.parse import quote
|
|||
import httpx
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.url_utils import SSRFError, assert_same_origin
|
||||
from litellm.constants import (
|
||||
AZURE_DOCUMENT_INTELLIGENCE_API_VERSION,
|
||||
AZURE_DOCUMENT_INTELLIGENCE_DEFAULT_DPI,
|
||||
|
|
@ -599,6 +600,16 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
|
|||
"Azure Document Intelligence returned 202 but no Operation-Location header found"
|
||||
)
|
||||
|
||||
# Reject cross-origin polling URLs — the auth headers
|
||||
# below would otherwise leak to whatever URL the upstream
|
||||
# (or an attacker-controlled upstream) returns. VERIA-51.
|
||||
try:
|
||||
assert_same_origin(operation_url, str(raw_response.request.url))
|
||||
except SSRFError as ssrf_err:
|
||||
raise ValueError(
|
||||
f"Azure Document Intelligence: rejected polling URL ({ssrf_err})"
|
||||
)
|
||||
|
||||
# Get headers for polling (need auth)
|
||||
poll_headers = {
|
||||
"Ocp-Apim-Subscription-Key": raw_response.request.headers.get(
|
||||
|
|
@ -711,6 +722,14 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
|
|||
"Azure Document Intelligence returned 202 but no Operation-Location header found"
|
||||
)
|
||||
|
||||
# Reject cross-origin polling URLs (see sync path). VERIA-51.
|
||||
try:
|
||||
assert_same_origin(operation_url, str(raw_response.request.url))
|
||||
except SSRFError as ssrf_err:
|
||||
raise ValueError(
|
||||
f"Azure Document Intelligence: rejected polling URL ({ssrf_err})"
|
||||
)
|
||||
|
||||
# Get headers for polling (need auth)
|
||||
poll_headers = {
|
||||
"Ocp-Apim-Subscription-Key": raw_response.request.headers.get(
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ import httpx
|
|||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.url_utils import SSRFError, assert_same_origin
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
AsyncHTTPHandler,
|
||||
HTTPHandler,
|
||||
|
|
@ -331,6 +332,17 @@ class BlackForestLabsImageEdit:
|
|||
message="No polling_url in BFL response",
|
||||
)
|
||||
|
||||
# Reject cross-origin polling URLs — the ``x-key`` auth header
|
||||
# would otherwise leak to whatever URL the upstream returns.
|
||||
# VERIA-51.
|
||||
try:
|
||||
assert_same_origin(polling_url, str(initial_response.request.url))
|
||||
except SSRFError as ssrf_err:
|
||||
raise BlackForestLabsError(
|
||||
status_code=502,
|
||||
message=f"Rejected polling URL: {ssrf_err}",
|
||||
)
|
||||
|
||||
# Get just the auth header for polling
|
||||
polling_headers = {"x-key": headers.get("x-key", "")}
|
||||
|
||||
|
|
@ -416,6 +428,17 @@ class BlackForestLabsImageEdit:
|
|||
message="No polling_url in BFL response",
|
||||
)
|
||||
|
||||
# Reject cross-origin polling URLs — the ``x-key`` auth header
|
||||
# would otherwise leak to whatever URL the upstream returns.
|
||||
# VERIA-51.
|
||||
try:
|
||||
assert_same_origin(polling_url, str(initial_response.request.url))
|
||||
except SSRFError as ssrf_err:
|
||||
raise BlackForestLabsError(
|
||||
status_code=502,
|
||||
message=f"Rejected polling URL: {ssrf_err}",
|
||||
)
|
||||
|
||||
# Get just the auth header for polling
|
||||
polling_headers = {"x-key": headers.get("x-key", "")}
|
||||
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ import httpx
|
|||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.url_utils import SSRFError, assert_same_origin
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
AsyncHTTPHandler,
|
||||
HTTPHandler,
|
||||
|
|
@ -317,6 +318,17 @@ class BlackForestLabsImageGeneration:
|
|||
message="No polling_url in BFL response",
|
||||
)
|
||||
|
||||
# Reject cross-origin polling URLs — the ``x-key`` auth header
|
||||
# would otherwise leak to whatever URL the upstream returns.
|
||||
# VERIA-51.
|
||||
try:
|
||||
assert_same_origin(polling_url, str(initial_response.request.url))
|
||||
except SSRFError as ssrf_err:
|
||||
raise BlackForestLabsError(
|
||||
status_code=502,
|
||||
message=f"Rejected polling URL: {ssrf_err}",
|
||||
)
|
||||
|
||||
# Get just the auth header for polling
|
||||
polling_headers = {"x-key": headers.get("x-key", "")}
|
||||
|
||||
|
|
@ -402,6 +414,17 @@ class BlackForestLabsImageGeneration:
|
|||
message="No polling_url in BFL response",
|
||||
)
|
||||
|
||||
# Reject cross-origin polling URLs — the ``x-key`` auth header
|
||||
# would otherwise leak to whatever URL the upstream returns.
|
||||
# VERIA-51.
|
||||
try:
|
||||
assert_same_origin(polling_url, str(initial_response.request.url))
|
||||
except SSRFError as ssrf_err:
|
||||
raise BlackForestLabsError(
|
||||
status_code=502,
|
||||
message=f"Rejected polling URL: {ssrf_err}",
|
||||
)
|
||||
|
||||
# Get just the auth header for polling
|
||||
polling_headers = {"x-key": headers.get("x-key", "")}
|
||||
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ Google AI Studio /batchEmbedContents Embeddings Endpoint
|
|||
"""
|
||||
|
||||
import json
|
||||
from typing import Any, Dict, Literal, Optional, Union
|
||||
from typing import Any, Dict, List, Literal, Optional, Tuple, Union
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -13,8 +13,8 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
HTTPHandler,
|
||||
get_async_httpx_client,
|
||||
)
|
||||
from litellm.types.llms.openai import EmbeddingInput
|
||||
from litellm.types.llms.vertex_ai import (
|
||||
GeminiEmbeddingInput,
|
||||
VertexAIBatchEmbeddingsRequestBody,
|
||||
VertexAIBatchEmbeddingsResponseObject,
|
||||
)
|
||||
|
|
@ -23,7 +23,6 @@ from litellm.types.utils import EmbeddingResponse
|
|||
from ..gemini.vertex_and_google_ai_studio_gemini import VertexLLM
|
||||
from .batch_embed_content_transformation import (
|
||||
_is_file_reference,
|
||||
_is_multimodal_input,
|
||||
process_embed_content_response,
|
||||
process_response,
|
||||
transform_openai_input_gemini_content,
|
||||
|
|
@ -32,9 +31,24 @@ from .batch_embed_content_transformation import (
|
|||
|
||||
|
||||
class GoogleBatchEmbeddings(VertexLLM):
|
||||
@staticmethod
|
||||
def _flatten_and_detect_file_refs(
|
||||
input: GeminiEmbeddingInput,
|
||||
) -> Tuple[List[str], bool]:
|
||||
"""Flatten nested input lists and detect file references."""
|
||||
input_list = [input] if isinstance(input, str) else input
|
||||
flat_elements = [
|
||||
e
|
||||
for item in input_list
|
||||
for e in (item if isinstance(item, list) else [item])
|
||||
if isinstance(e, str)
|
||||
]
|
||||
has_file_refs = any(_is_file_reference(e) for e in flat_elements)
|
||||
return flat_elements, has_file_refs
|
||||
|
||||
def _resolve_file_references(
|
||||
self,
|
||||
input: EmbeddingInput,
|
||||
input: GeminiEmbeddingInput,
|
||||
api_key: str,
|
||||
sync_handler: HTTPHandler,
|
||||
) -> Dict[str, Dict[str, str]]:
|
||||
|
|
@ -42,7 +56,7 @@ class GoogleBatchEmbeddings(VertexLLM):
|
|||
Resolve Gemini file references (files/...) to get mime_type and uri.
|
||||
|
||||
Args:
|
||||
input: EmbeddingInput that may contain file references
|
||||
input: GeminiEmbeddingInput that may contain file references
|
||||
api_key: Gemini API key
|
||||
sync_handler: HTTP client
|
||||
|
||||
|
|
@ -73,7 +87,7 @@ class GoogleBatchEmbeddings(VertexLLM):
|
|||
|
||||
async def _async_resolve_file_references(
|
||||
self,
|
||||
input: EmbeddingInput,
|
||||
input: GeminiEmbeddingInput,
|
||||
api_key: str,
|
||||
async_handler: AsyncHTTPHandler,
|
||||
) -> Dict[str, Dict[str, str]]:
|
||||
|
|
@ -81,7 +95,7 @@ class GoogleBatchEmbeddings(VertexLLM):
|
|||
Async version of _resolve_file_references.
|
||||
|
||||
Args:
|
||||
input: EmbeddingInput that may contain file references
|
||||
input: GeminiEmbeddingInput that may contain file references
|
||||
api_key: Gemini API key
|
||||
async_handler: Async HTTP client
|
||||
|
||||
|
|
@ -110,10 +124,10 @@ class GoogleBatchEmbeddings(VertexLLM):
|
|||
|
||||
return resolved_files
|
||||
|
||||
def batch_embeddings(
|
||||
def batch_embeddings( # noqa: PLR0915
|
||||
self,
|
||||
model: str,
|
||||
input: EmbeddingInput,
|
||||
input: GeminiEmbeddingInput,
|
||||
print_verbose,
|
||||
model_response: EmbeddingResponse,
|
||||
custom_llm_provider: Literal["gemini", "vertex_ai"],
|
||||
|
|
@ -151,8 +165,7 @@ class GoogleBatchEmbeddings(VertexLLM):
|
|||
|
||||
optional_params = optional_params or {}
|
||||
|
||||
is_multimodal = _is_multimodal_input(input)
|
||||
use_embed_content = is_multimodal or (custom_llm_provider == "vertex_ai")
|
||||
use_embed_content = custom_llm_provider == "vertex_ai"
|
||||
mode: Literal["embedding", "batch_embedding"]
|
||||
if use_embed_content:
|
||||
mode = "embedding"
|
||||
|
|
@ -215,8 +228,22 @@ class GoogleBatchEmbeddings(VertexLLM):
|
|||
resolved_files=resolved_files,
|
||||
)
|
||||
else:
|
||||
flat_elements, has_file_refs = self._flatten_and_detect_file_refs(input)
|
||||
if has_file_refs and not api_key:
|
||||
raise ValueError(
|
||||
"An API key is required to resolve Gemini file references (files/...). "
|
||||
"Pass api_key= or set GEMINI_API_KEY."
|
||||
)
|
||||
resolved_files = {}
|
||||
if api_key and has_file_refs:
|
||||
resolved_files = self._resolve_file_references(
|
||||
input=flat_elements, api_key=api_key, sync_handler=sync_handler
|
||||
)
|
||||
request_data = transform_openai_input_gemini_content(
|
||||
input=input, model=model, optional_params=optional_params
|
||||
input=input,
|
||||
model=model,
|
||||
optional_params=optional_params,
|
||||
resolved_files=resolved_files,
|
||||
)
|
||||
|
||||
## LOGGING
|
||||
|
|
@ -264,7 +291,7 @@ class GoogleBatchEmbeddings(VertexLLM):
|
|||
url: str,
|
||||
data: Optional[Union[VertexAIBatchEmbeddingsRequestBody, dict]],
|
||||
model_response: EmbeddingResponse,
|
||||
input: EmbeddingInput,
|
||||
input: GeminiEmbeddingInput,
|
||||
timeout: Optional[Union[float, httpx.Timeout]],
|
||||
headers={},
|
||||
client: Optional[AsyncHTTPHandler] = None,
|
||||
|
|
@ -303,8 +330,22 @@ class GoogleBatchEmbeddings(VertexLLM):
|
|||
resolved_files=resolved_files,
|
||||
)
|
||||
else:
|
||||
flat_elements, has_file_refs = self._flatten_and_detect_file_refs(input)
|
||||
if has_file_refs and not api_key:
|
||||
raise ValueError(
|
||||
"An API key is required to resolve Gemini file references (files/...). "
|
||||
"Pass api_key= or set GEMINI_API_KEY."
|
||||
)
|
||||
resolved_files = {}
|
||||
if api_key and has_file_refs:
|
||||
resolved_files = await self._async_resolve_file_references(
|
||||
input=flat_elements, api_key=api_key, async_handler=async_handler
|
||||
)
|
||||
data = transform_openai_input_gemini_content(
|
||||
input=input, model=model, optional_params=optional_params or {}
|
||||
input=input,
|
||||
model=model,
|
||||
optional_params=optional_params or {},
|
||||
resolved_files=resolved_files,
|
||||
)
|
||||
|
||||
## LOGGING
|
||||
|
|
|
|||
|
|
@ -6,12 +6,12 @@ Why separate file? Make it easy to see how transformation works
|
|||
|
||||
from typing import Dict, List, Optional, Tuple
|
||||
|
||||
from litellm.types.llms.openai import EmbeddingInput
|
||||
from litellm.types.llms.vertex_ai import (
|
||||
BlobType,
|
||||
ContentType,
|
||||
EmbedContentRequest,
|
||||
FileDataType,
|
||||
GeminiEmbeddingInput,
|
||||
PartType,
|
||||
VertexAIBatchEmbeddingsRequestBody,
|
||||
VertexAIBatchEmbeddingsResponseObject,
|
||||
|
|
@ -114,33 +114,77 @@ def _parse_data_url(data_url: str) -> Tuple[str, str]:
|
|||
return media_type, base64_data
|
||||
|
||||
|
||||
def _is_multimodal_input(input: EmbeddingInput) -> bool:
|
||||
def _is_multimodal_input(input: GeminiEmbeddingInput) -> bool:
|
||||
"""
|
||||
Check if the input contains multimodal data (data URIs, file references, or GCS URLs).
|
||||
Check if the input contains multimodal data (data URIs, file references,
|
||||
GCS URLs, or nested lists for combined embeddings).
|
||||
|
||||
Args:
|
||||
input: EmbeddingInput (str or List[str])
|
||||
input: GeminiEmbeddingInput — str, List[str], or List[List[str]] for combined embeddings
|
||||
|
||||
Returns:
|
||||
bool: True if any element is a data URI, file reference, or GCS URL
|
||||
bool: True if any element is multimodal or a nested list
|
||||
"""
|
||||
if isinstance(input, str):
|
||||
input_list = [input]
|
||||
else:
|
||||
input_list = input
|
||||
return _is_multimodal_element(input)
|
||||
|
||||
for element in input_list:
|
||||
if isinstance(element, str):
|
||||
if element.startswith("data:") and ";base64," in element:
|
||||
return True
|
||||
if _is_file_reference(element):
|
||||
return True
|
||||
if _is_gcs_url(element):
|
||||
for element in input:
|
||||
if isinstance(element, list):
|
||||
if any(
|
||||
_is_multimodal_element(sub) for sub in element if isinstance(sub, str)
|
||||
):
|
||||
return True
|
||||
elif isinstance(element, str) and _is_multimodal_element(element):
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
|
||||
def _is_multimodal_element(element: str) -> bool:
|
||||
"""Check if a single string element is multimodal."""
|
||||
if element.startswith("data:") and ";base64," in element:
|
||||
return True
|
||||
if _is_file_reference(element):
|
||||
return True
|
||||
if _is_gcs_url(element):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def _build_part_for_input(
|
||||
element: str,
|
||||
resolved_files: Optional[Dict[str, Dict[str, str]]] = None,
|
||||
) -> PartType:
|
||||
"""
|
||||
Build a single PartType for an input element, handling text, data URIs,
|
||||
file references, and GCS URLs.
|
||||
"""
|
||||
resolved_files = resolved_files or {}
|
||||
|
||||
if element.startswith("data:") and ";base64," in element:
|
||||
mime_type, base64_data = _parse_data_url(element)
|
||||
blob: BlobType = {"mime_type": mime_type, "data": base64_data}
|
||||
return PartType(inline_data=blob)
|
||||
elif _is_gcs_url(element):
|
||||
mime_type = _infer_mime_type_from_gcs_url(element)
|
||||
file_data: FileDataType = {
|
||||
"mime_type": mime_type,
|
||||
"file_uri": element,
|
||||
}
|
||||
return PartType(file_data=file_data)
|
||||
elif _is_file_reference(element):
|
||||
if element not in resolved_files:
|
||||
raise ValueError(f"File reference {element} not resolved")
|
||||
file_info = resolved_files[element]
|
||||
file_data_ref: FileDataType = {
|
||||
"mime_type": file_info["mime_type"],
|
||||
"file_uri": file_info["uri"],
|
||||
}
|
||||
return PartType(file_data=file_data_ref)
|
||||
else:
|
||||
return PartType(text=element)
|
||||
|
||||
|
||||
_SUPPORTED_EMBED_PARAMS = {"outputDimensionality", "taskType", "title"}
|
||||
|
||||
|
||||
|
|
@ -155,37 +199,60 @@ def _filter_embed_params(optional_params: dict) -> dict:
|
|||
|
||||
|
||||
def transform_openai_input_gemini_content(
|
||||
input: EmbeddingInput, model: str, optional_params: dict
|
||||
input: GeminiEmbeddingInput,
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
resolved_files: Optional[Dict[str, Dict[str, str]]] = None,
|
||||
) -> VertexAIBatchEmbeddingsRequestBody:
|
||||
"""
|
||||
The content to embed. Only the parts.text fields will be counted.
|
||||
Transform OpenAI embedding input to Gemini batchEmbedContents format.
|
||||
|
||||
Each input element becomes a separate EmbedContentRequest, supporting
|
||||
text, data URIs, file references, and GCS URLs.
|
||||
|
||||
If an element is a list (nested input), all sub-elements are combined
|
||||
into a single content with multiple parts, producing one combined
|
||||
embedding for the group.
|
||||
|
||||
Examples:
|
||||
input=["text", "image"] → 2 separate embeddings
|
||||
input=[["text", "image"]] → 1 combined embedding
|
||||
input=[["text", "image"], "x"] → 2 embeddings (1 combined + 1 separate)
|
||||
"""
|
||||
gemini_model_name = "models/{}".format(model)
|
||||
|
||||
gemini_params = _filter_embed_params(optional_params)
|
||||
|
||||
input_list = [input] if isinstance(input, str) else input
|
||||
requests: List[EmbedContentRequest] = []
|
||||
if isinstance(input, str):
|
||||
|
||||
for element in input_list:
|
||||
if isinstance(element, list):
|
||||
if not element:
|
||||
raise ValueError("Nested input list must not be empty")
|
||||
for sub in element:
|
||||
if not isinstance(sub, str):
|
||||
raise ValueError(
|
||||
f"Elements inside a nested input list must be strings, got {type(sub)}"
|
||||
)
|
||||
parts = [
|
||||
_build_part_for_input(sub, resolved_files=resolved_files)
|
||||
for sub in element
|
||||
]
|
||||
else:
|
||||
parts = [_build_part_for_input(element, resolved_files=resolved_files)]
|
||||
request = EmbedContentRequest(
|
||||
model=gemini_model_name,
|
||||
content=ContentType(parts=[PartType(text=input)]),
|
||||
content=ContentType(parts=parts),
|
||||
**gemini_params,
|
||||
)
|
||||
requests.append(request)
|
||||
else:
|
||||
for i in input:
|
||||
request = EmbedContentRequest(
|
||||
model=gemini_model_name,
|
||||
content=ContentType(parts=[PartType(text=i)]),
|
||||
**gemini_params,
|
||||
)
|
||||
requests.append(request)
|
||||
|
||||
return VertexAIBatchEmbeddingsRequestBody(requests=requests)
|
||||
|
||||
|
||||
def transform_openai_input_gemini_embed_content(
|
||||
input: EmbeddingInput,
|
||||
input: GeminiEmbeddingInput,
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
resolved_files: Optional[Dict[str, Dict[str, str]]] = None,
|
||||
|
|
@ -194,7 +261,7 @@ def transform_openai_input_gemini_embed_content(
|
|||
Transform OpenAI embedding input to Gemini embedContent format (multimodal).
|
||||
|
||||
Args:
|
||||
input: EmbeddingInput (str or List[str]) with text, data URIs, or file references
|
||||
input: GeminiEmbeddingInput with text, data URIs, or file references
|
||||
model: Model name
|
||||
optional_params: Additional parameters (taskType, outputDimensionality, etc.)
|
||||
resolved_files: Dict mapping file names (files/abc) to {mime_type, uri}
|
||||
|
|
@ -210,31 +277,14 @@ def transform_openai_input_gemini_embed_content(
|
|||
parts: List[PartType] = []
|
||||
|
||||
for element in input_list:
|
||||
if isinstance(element, list):
|
||||
raise ValueError(
|
||||
"Nested (combined) embeddings are not supported on the embedContent path. "
|
||||
"Use the batchEmbedContents path or pass a flat list instead."
|
||||
)
|
||||
if not isinstance(element, str):
|
||||
raise ValueError(f"Unsupported input type: {type(element)}")
|
||||
|
||||
if element.startswith("data:") and ";base64," in element:
|
||||
mime_type, base64_data = _parse_data_url(element)
|
||||
blob: BlobType = {"mime_type": mime_type, "data": base64_data}
|
||||
parts.append(PartType(inline_data=blob))
|
||||
elif _is_gcs_url(element):
|
||||
mime_type = _infer_mime_type_from_gcs_url(element)
|
||||
file_data: FileDataType = {
|
||||
"mime_type": mime_type,
|
||||
"file_uri": element,
|
||||
}
|
||||
parts.append(PartType(file_data=file_data))
|
||||
elif _is_file_reference(element):
|
||||
if element not in resolved_files:
|
||||
raise ValueError(f"File reference {element} not resolved")
|
||||
file_info = resolved_files[element]
|
||||
file_data_ref: FileDataType = {
|
||||
"mime_type": file_info["mime_type"],
|
||||
"file_uri": file_info["uri"],
|
||||
}
|
||||
parts.append(PartType(file_data=file_data_ref))
|
||||
else:
|
||||
parts.append(PartType(text=element))
|
||||
parts.append(_build_part_for_input(element, resolved_files=resolved_files))
|
||||
|
||||
request_body: dict = {
|
||||
"content": ContentType(parts=parts),
|
||||
|
|
@ -245,7 +295,7 @@ def transform_openai_input_gemini_embed_content(
|
|||
|
||||
|
||||
def process_embed_content_response(
|
||||
input: EmbeddingInput,
|
||||
input: GeminiEmbeddingInput,
|
||||
model_response: EmbeddingResponse,
|
||||
model: str,
|
||||
response_json: dict,
|
||||
|
|
@ -291,7 +341,7 @@ def process_embed_content_response(
|
|||
|
||||
|
||||
def process_response(
|
||||
input: EmbeddingInput,
|
||||
input: GeminiEmbeddingInput,
|
||||
model_response: EmbeddingResponse,
|
||||
model: str,
|
||||
_predictions: VertexAIBatchEmbeddingsResponseObject,
|
||||
|
|
@ -308,8 +358,29 @@ def process_response(
|
|||
model_response.data = openai_embeddings
|
||||
model_response.model = model
|
||||
|
||||
input_text = get_formatted_prompt(data={"input": input}, call_type="embedding")
|
||||
prompt_tokens = token_counter(model=model, text=input_text)
|
||||
has_nested = isinstance(input, list) and any(isinstance(e, list) for e in input)
|
||||
if _is_multimodal_input(input) or has_nested:
|
||||
input_list = input if isinstance(input, list) else [input]
|
||||
text_elements: List[str] = []
|
||||
for e in input_list:
|
||||
if isinstance(e, list):
|
||||
text_elements.extend(
|
||||
sub
|
||||
for sub in e
|
||||
if isinstance(sub, str) and not _is_multimodal_element(sub)
|
||||
)
|
||||
elif isinstance(e, str) and not _is_multimodal_element(e):
|
||||
text_elements.append(e)
|
||||
if text_elements:
|
||||
input_text = get_formatted_prompt(
|
||||
data={"input": text_elements}, call_type="embedding"
|
||||
)
|
||||
prompt_tokens = token_counter(model=model, text=input_text)
|
||||
else:
|
||||
prompt_tokens = 0
|
||||
else:
|
||||
input_text = get_formatted_prompt(data={"input": input}, call_type="embedding")
|
||||
prompt_tokens = token_counter(model=model, text=input_text)
|
||||
model_response.usage = Usage(
|
||||
prompt_tokens=prompt_tokens, total_tokens=prompt_tokens
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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.",
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -11,9 +11,19 @@ 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_METHOD_SUFFIXES = {
|
||||
"delete",
|
||||
"get",
|
||||
"head",
|
||||
"options",
|
||||
"patch",
|
||||
"post",
|
||||
"put",
|
||||
"trace",
|
||||
}
|
||||
|
||||
|
||||
def _stabilize_multi_method_route_ids(routes) -> None:
|
||||
|
|
@ -39,13 +49,46 @@ def load_snapshot() -> Optional[Dict[str, Dict]]:
|
|||
return None
|
||||
|
||||
|
||||
def _normalize_operation_ids(paths: Dict[str, Dict]) -> None:
|
||||
"""Make FastAPI-generated operation IDs stable for multi-method routes.
|
||||
|
||||
FastAPI derives the default operation ID suffix from the first item in the
|
||||
route's methods set. For routes registered with several HTTP methods, that
|
||||
set iteration order can vary between processes, which makes the snapshot
|
||||
drift even when no routes changed.
|
||||
"""
|
||||
for path_ops in paths.values():
|
||||
if not isinstance(path_ops, dict):
|
||||
continue
|
||||
|
||||
methods = {method for method in path_ops if method in HTTP_METHODS}
|
||||
if not methods:
|
||||
continue
|
||||
|
||||
for method, operation in path_ops.items():
|
||||
if method not in HTTP_METHODS or not isinstance(operation, dict):
|
||||
continue
|
||||
|
||||
operation_id = operation.get("operationId")
|
||||
if not isinstance(operation_id, str):
|
||||
continue
|
||||
|
||||
for suffix in methods:
|
||||
suffix_token = f"_{suffix}"
|
||||
if operation_id.endswith(suffix_token):
|
||||
operation["operationId"] = (
|
||||
operation_id[: -len(suffix_token)] + f"_{method}"
|
||||
)
|
||||
break
|
||||
|
||||
|
||||
def generate_snapshot() -> Dict[str, Dict]:
|
||||
import importlib
|
||||
|
||||
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:
|
||||
|
|
@ -57,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
|
||||
|
|
@ -67,13 +111,24 @@ def generate_snapshot() -> Dict[str, Dict]:
|
|||
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 full.get("paths", {}).values():
|
||||
for op in path_ops.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": full.get("paths", {}),
|
||||
"paths": paths,
|
||||
"components": {"schemas": full.get("components", {}).get("schemas", {})},
|
||||
}
|
||||
return fragments
|
||||
|
|
|
|||
|
|
@ -668,6 +668,8 @@ class LiteLLMRoutes(enum.Enum):
|
|||
"/models/{model_id}",
|
||||
"/guardrails/list",
|
||||
"/v2/guardrails/list",
|
||||
"/project/list",
|
||||
"/project/info",
|
||||
]
|
||||
+ spend_tracking_routes
|
||||
+ key_management_routes
|
||||
|
|
@ -692,6 +694,9 @@ class LiteLLMRoutes(enum.Enum):
|
|||
"/model/{model_id}/update",
|
||||
"/prompt/list",
|
||||
"/prompt/info",
|
||||
# Project read routes - endpoint scopes results to caller's teams (non-admin)
|
||||
"/project/list",
|
||||
"/project/info",
|
||||
# Invitation routes - org/team admins checked in endpoint via _user_has_admin_privileges
|
||||
"/invitation/new",
|
||||
"/invitation/delete",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
@ -167,6 +167,81 @@ def _allow_model_level_clientside_configurable_parameters(
|
|||
)
|
||||
|
||||
|
||||
# Config dicts whose entries are spread as ``**dict`` into outbound LLM
|
||||
# API calls. ``litellm_embedding_config`` is consumed by the Milvus
|
||||
# vector store transformer; future nested-config keys with the same
|
||||
# threat shape should be added here.
|
||||
_NESTED_CONFIG_KEYS: Tuple[str, ...] = ("litellm_embedding_config",)
|
||||
|
||||
# Banned root-level params. Same list applies to every entry in
|
||||
# ``_NESTED_CONFIG_KEYS`` because those dicts get spread as ``**kwargs``
|
||||
# into the same outbound calls.
|
||||
_BANNED_REQUEST_BODY_PARAMS: Tuple[str, ...] = (
|
||||
"api_base",
|
||||
"base_url",
|
||||
"user_config",
|
||||
"aws_sts_endpoint",
|
||||
"aws_web_identity_token",
|
||||
"aws_role_name",
|
||||
"vertex_credentials",
|
||||
# Endpoint-targeting fields that retarget the outbound request or
|
||||
# an observability callback. An attacker-controlled value either
|
||||
# exfiltrates the request payload (incl. messages + admin-set
|
||||
# tokens) to the attacker's host, or coerces the proxy into
|
||||
# authenticating against the attacker's host with admin secrets.
|
||||
"aws_bedrock_runtime_endpoint",
|
||||
"langsmith_base_url",
|
||||
"langfuse_host",
|
||||
"posthog_host",
|
||||
"braintrust_host",
|
||||
"slack_webhook_url",
|
||||
# Provider-specific endpoint overrides that flow into the outbound
|
||||
# request via ``optional_params``. Same threat as ``api_base``:
|
||||
# ``s3_endpoint_url`` redirects Bedrock file uploads to attacker
|
||||
# S3; ``sagemaker_base_url`` redirects all SageMaker traffic;
|
||||
# ``deployment_url`` redirects SAP deployments.
|
||||
"s3_endpoint_url",
|
||||
"sagemaker_base_url",
|
||||
"deployment_url",
|
||||
)
|
||||
|
||||
|
||||
def _check_banned_params(
|
||||
body: dict,
|
||||
general_settings: dict,
|
||||
llm_router: Optional[Router],
|
||||
model: str,
|
||||
) -> None:
|
||||
"""Raise ``ValueError`` if ``body`` carries a banned param without admin opt-in.
|
||||
|
||||
Shared between the root-level check and the nested-config check so a
|
||||
new banned param only needs to be added in one place.
|
||||
"""
|
||||
for param in _BANNED_REQUEST_BODY_PARAMS:
|
||||
if param not in body:
|
||||
continue
|
||||
if general_settings.get("allow_client_side_credentials") is True:
|
||||
return
|
||||
if (
|
||||
_allow_model_level_clientside_configurable_parameters(
|
||||
model=model,
|
||||
param=param,
|
||||
request_body_value=body[param],
|
||||
llm_router=llm_router,
|
||||
)
|
||||
is True
|
||||
):
|
||||
return
|
||||
raise ValueError(
|
||||
f"Rejected Request: {param} is not allowed in request body. "
|
||||
"Clientside passthrough requires explicit admin opt-in via "
|
||||
"either `general_settings.allow_client_side_credentials = true` "
|
||||
"(proxy-wide) or `configurable_clientside_auth_params` on the "
|
||||
"deployment in your proxy config.yaml. "
|
||||
"Relevant Issue: https://huntr.com/bounties/4001e1a2-7b7a-4776-a3ae-e6692ec3d997",
|
||||
)
|
||||
|
||||
|
||||
def is_request_body_safe(
|
||||
request_body: dict, general_settings: dict, llm_router: Optional[Router], model: str
|
||||
) -> bool:
|
||||
|
|
@ -175,72 +250,31 @@ def is_request_body_safe(
|
|||
|
||||
A malicious user can set the api_base to their own domain and invoke POST /chat/completions to intercept and steal the OpenAI API key.
|
||||
Relevant issue: https://huntr.com/bounties/4001e1a2-7b7a-4776-a3ae-e6692ec3d997
|
||||
|
||||
The blocklist is enforced unconditionally. Legitimate clientside
|
||||
credential / endpoint passthrough goes through one of the two
|
||||
explicit admin opt-ins (``general_settings.allow_client_side_credentials``
|
||||
proxy-wide or ``configurable_clientside_auth_params`` per deployment).
|
||||
Historically there was a third, *implicit*, *caller-controlled* path:
|
||||
``check_complete_credentials`` returned True when the caller supplied
|
||||
any non-empty ``api_key``, which made the entire blocklist a no-op.
|
||||
That bypass turned every missing entry on the blocklist into an
|
||||
exploitable SSRF / credential-exfil hole — see GHSA-jh89-88fc-qrfp,
|
||||
GHSA-3frq-6r6h-7j64, and the chain of veria-admin findings (Dv_m860l,
|
||||
b_yRJeQ5, stN90yjP, LBlyOAc8, U2TD78kg). Removed: the blocklist now
|
||||
has a single, predictable failure mode for missing entries (a 400),
|
||||
not a credential leak.
|
||||
|
||||
Iterative single-level descent into ``_NESTED_CONFIG_KEYS`` (rather
|
||||
than recursion) covers nested-config attacks like Milvus's
|
||||
``litellm_embedding_config.api_base`` (VERIA-6) without exposing a
|
||||
recursion-depth DoS surface.
|
||||
"""
|
||||
banned_params = [
|
||||
"api_base",
|
||||
"base_url",
|
||||
"user_config",
|
||||
"aws_sts_endpoint",
|
||||
"aws_web_identity_token",
|
||||
"aws_role_name",
|
||||
"vertex_credentials",
|
||||
# Endpoint-targeting fields that retarget the outbound request or
|
||||
# an observability callback. An attacker-controlled value either
|
||||
# exfiltrates the request payload (incl. messages + admin-set
|
||||
# tokens) to the attacker's host, or coerces the proxy into
|
||||
# authenticating against the attacker's host with admin secrets.
|
||||
"aws_bedrock_runtime_endpoint",
|
||||
"langsmith_base_url",
|
||||
"langfuse_host",
|
||||
"posthog_host",
|
||||
"braintrust_host",
|
||||
"slack_webhook_url",
|
||||
# Provider-specific endpoint overrides that flow into the outbound
|
||||
# request via ``optional_params``. Same threat as ``api_base``:
|
||||
# ``s3_endpoint_url`` redirects Bedrock file uploads to attacker
|
||||
# S3; ``sagemaker_base_url`` redirects all SageMaker traffic;
|
||||
# ``deployment_url`` redirects SAP deployments.
|
||||
"s3_endpoint_url",
|
||||
"sagemaker_base_url",
|
||||
"deployment_url",
|
||||
]
|
||||
|
||||
# The blocklist is enforced unconditionally. Legitimate clientside
|
||||
# credential / endpoint passthrough goes through one of the two
|
||||
# explicit admin opt-ins (``general_settings.allow_client_side_credentials``
|
||||
# proxy-wide or ``configurable_clientside_auth_params`` per deployment).
|
||||
# Historically there was a third, *implicit*, *caller-controlled* path:
|
||||
# ``check_complete_credentials`` returned True when the caller supplied
|
||||
# any non-empty ``api_key``, which made the entire blocklist a no-op.
|
||||
# That bypass turned every missing entry on the blocklist into an
|
||||
# exploitable SSRF / credential-exfil hole — see GHSA-jh89-88fc-qrfp,
|
||||
# GHSA-3frq-6r6h-7j64, and the chain of veria-admin findings (Dv_m860l,
|
||||
# b_yRJeQ5, stN90yjP, LBlyOAc8, U2TD78kg). Removed: the blocklist now
|
||||
# has a single, predictable failure mode for missing entries (a 400),
|
||||
# not a credential leak.
|
||||
for param in banned_params:
|
||||
if param in request_body:
|
||||
if general_settings.get("allow_client_side_credentials") is True:
|
||||
return True
|
||||
elif (
|
||||
_allow_model_level_clientside_configurable_parameters(
|
||||
model=model,
|
||||
param=param,
|
||||
request_body_value=request_body[param],
|
||||
llm_router=llm_router,
|
||||
)
|
||||
is True
|
||||
):
|
||||
return True
|
||||
raise ValueError(
|
||||
f"Rejected Request: {param} is not allowed in request body. "
|
||||
"Clientside passthrough requires explicit admin opt-in via "
|
||||
"either `general_settings.allow_client_side_credentials = true` "
|
||||
"(proxy-wide) or `configurable_clientside_auth_params` on the "
|
||||
"deployment in your proxy config.yaml. "
|
||||
"Relevant Issue: https://huntr.com/bounties/4001e1a2-7b7a-4776-a3ae-e6692ec3d997",
|
||||
)
|
||||
|
||||
_check_banned_params(request_body, general_settings, llm_router, model)
|
||||
for nested_key in _NESTED_CONFIG_KEYS:
|
||||
nested = request_body.get(nested_key)
|
||||
if isinstance(nested, dict):
|
||||
_check_banned_params(nested, general_settings, llm_router, model)
|
||||
return True
|
||||
|
||||
|
||||
|
|
@ -942,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:
|
||||
|
|
@ -2199,6 +2241,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:
|
||||
|
|
@ -2217,7 +2260,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),
|
||||
|
|
@ -2304,11 +2351,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
|
||||
|
|
@ -2316,13 +2369,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
|
||||
|
|
@ -2330,14 +2384,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
|
||||
|
|
|
|||
|
|
@ -53,12 +53,16 @@ def clear_token() -> None:
|
|||
os.remove(token_file)
|
||||
|
||||
|
||||
def get_stored_api_key() -> Optional[str]:
|
||||
"""Get the stored API key from token file"""
|
||||
# Use the SDK-level utility
|
||||
def get_stored_api_key(expected_base_url: Optional[str] = None) -> Optional[str]:
|
||||
"""Get the stored API key from token file.
|
||||
|
||||
If expected_base_url is provided, the key is only returned when it was
|
||||
originally issued for that URL. This prevents credential leakage when the
|
||||
CLI is pointed at a different (possibly malicious) server.
|
||||
"""
|
||||
from litellm.litellm_core_utils.cli_token_utils import get_litellm_gateway_api_key
|
||||
|
||||
return get_litellm_gateway_api_key()
|
||||
return get_litellm_gateway_api_key(expected_base_url=expected_base_url)
|
||||
|
||||
|
||||
# Team selection utilities
|
||||
|
|
@ -572,9 +576,11 @@ def login(ctx: click.Context):
|
|||
api_key = auth_result["api_key"]
|
||||
user_id = auth_result["user_id"]
|
||||
|
||||
# Save token data (simplified for CLI - we just need the key)
|
||||
# Save token data. base_url is stored so we can verify origin
|
||||
# before reusing the key on a subsequent CLI invocation.
|
||||
save_token(
|
||||
{
|
||||
"base_url": base_url.rstrip("/"),
|
||||
"key": api_key,
|
||||
"user_id": user_id or "cli-user",
|
||||
"user_email": "unknown",
|
||||
|
|
|
|||
|
|
@ -74,9 +74,10 @@ def cli(ctx: click.Context, base_url: str, api_key: Optional[str]) -> None:
|
|||
"""LiteLLM Proxy CLI - Manage your LiteLLM proxy server"""
|
||||
ctx.ensure_object(dict)
|
||||
|
||||
# If no API key provided via flag or environment variable, try to load from saved token
|
||||
# If no API key provided via flag or environment variable, try to load from saved token.
|
||||
# Pass base_url so we only use the stored key when it was issued for this server.
|
||||
if api_key is None:
|
||||
api_key = get_stored_api_key()
|
||||
api_key = get_stored_api_key(expected_base_url=base_url)
|
||||
|
||||
ctx.obj["base_url"] = base_url
|
||||
ctx.obj["api_key"] = api_key
|
||||
|
|
|
|||
|
|
@ -28,12 +28,17 @@ class Client:
|
|||
api_key (Optional[str]): API key for authentication. If provided, it will be sent as a Bearer token.
|
||||
timeout: Request timeout in seconds (default: 30)
|
||||
"""
|
||||
self._base_url = base_url.rstrip("/") # Remove trailing slash if present
|
||||
self._api_key = get_litellm_gateway_api_key() or api_key
|
||||
self._base_url = base_url.rstrip("/")
|
||||
# Only use the stored CLI key when it was issued for this server.
|
||||
self._api_key = api_key or get_litellm_gateway_api_key(
|
||||
expected_base_url=self._base_url
|
||||
)
|
||||
|
||||
# Initialize resource clients
|
||||
|
||||
self.http = HTTPClient(base_url=base_url, api_key=api_key, timeout=timeout)
|
||||
self.http = HTTPClient(
|
||||
base_url=base_url, api_key=self._api_key, timeout=timeout
|
||||
)
|
||||
self.models = ModelsManagementClient(
|
||||
base_url=self._base_url, api_key=self._api_key
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -225,10 +225,10 @@ class ToolPermissionGuardrail(CustomGuardrail):
|
|||
|
||||
def _parse_tool_call_arguments(
|
||||
self, tool_call: ChatCompletionMessageToolCall
|
||||
) -> Dict[str, Any]:
|
||||
) -> tuple[Optional[Dict[str, Any]], Optional[str]]:
|
||||
arguments = getattr(tool_call.function, "arguments", None)
|
||||
if not arguments:
|
||||
return {}
|
||||
return None, "missing arguments"
|
||||
|
||||
parsed_arguments: Any = {}
|
||||
try:
|
||||
|
|
@ -236,22 +236,24 @@ class ToolPermissionGuardrail(CustomGuardrail):
|
|||
parsed_arguments = json.loads(arguments)
|
||||
elif isinstance(arguments, dict):
|
||||
parsed_arguments = arguments
|
||||
except json.JSONDecodeError as exc:
|
||||
else:
|
||||
return None, "arguments must be a JSON object"
|
||||
except (json.JSONDecodeError, TypeError) as exc:
|
||||
verbose_proxy_logger.warning(
|
||||
"Tool Permission Guardrail: Failed to decode arguments for tool %s: %s",
|
||||
tool_call.function.name,
|
||||
exc,
|
||||
)
|
||||
return {}
|
||||
return None, "arguments could not be parsed"
|
||||
|
||||
if isinstance(parsed_arguments, dict):
|
||||
return parsed_arguments
|
||||
return parsed_arguments, None
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
"Tool Permission Guardrail: Ignoring non-dict arguments for tool %s",
|
||||
"Tool Permission Guardrail: Rejecting non-dict arguments for tool %s",
|
||||
tool_call.function.name,
|
||||
)
|
||||
return {}
|
||||
return None, "arguments must be a JSON object"
|
||||
|
||||
def _collect_argument_paths(
|
||||
self,
|
||||
|
|
@ -331,10 +333,21 @@ class ToolPermissionGuardrail(CustomGuardrail):
|
|||
continue
|
||||
|
||||
if rule.allowed_param_patterns and should_check_params:
|
||||
arguments = self._parse_tool_call_arguments(tool_call)
|
||||
arguments, parse_error = self._parse_tool_call_arguments(tool_call)
|
||||
if parse_error:
|
||||
default_message = f"Tool '{tool_identifier}' {parse_error} required by rule '{rule.id}'"
|
||||
message = self.render_violation_message(
|
||||
default=default_message,
|
||||
context={"tool_name": tool_identifier, "rule_id": rule.id},
|
||||
)
|
||||
return False, rule.id, message
|
||||
if not arguments:
|
||||
last_pattern_failure_msg = f"Tool '{tool_identifier}' is missing arguments required by rule '{rule.id}'"
|
||||
continue
|
||||
default_message = f"Tool '{tool_identifier}' is missing arguments required by rule '{rule.id}'"
|
||||
message = self.render_violation_message(
|
||||
default=default_message,
|
||||
context={"tool_name": tool_identifier, "rule_id": rule.id},
|
||||
)
|
||||
return False, rule.id, message
|
||||
|
||||
patterns_match, failure_message = self._patterns_match_for_rule(
|
||||
arguments=arguments,
|
||||
|
|
@ -365,6 +378,33 @@ class ToolPermissionGuardrail(CustomGuardrail):
|
|||
)
|
||||
return is_allowed, None, message
|
||||
|
||||
@staticmethod
|
||||
def _get_mapping_value(item: Any, key: str) -> Any:
|
||||
if isinstance(item, dict):
|
||||
return item.get(key)
|
||||
return getattr(item, key, None)
|
||||
|
||||
@staticmethod
|
||||
def _legacy_function_call_id(choice_index: int) -> str:
|
||||
return f"legacy_function_call_{choice_index}"
|
||||
|
||||
def _legacy_function_call_to_tool_call(
|
||||
self, function_call: Any, choice_index: int
|
||||
) -> Optional[ChatCompletionMessageToolCall]:
|
||||
if function_call is None:
|
||||
return None
|
||||
|
||||
function_name = self._get_mapping_value(function_call, "name")
|
||||
arguments = self._get_mapping_value(function_call, "arguments") or ""
|
||||
if not function_name:
|
||||
return None
|
||||
|
||||
return ChatCompletionMessageToolCall(
|
||||
id=self._legacy_function_call_id(choice_index),
|
||||
type="function",
|
||||
function={"name": function_name, "arguments": arguments},
|
||||
)
|
||||
|
||||
def _extract_tool_calls_from_response(
|
||||
self, response: ModelResponse
|
||||
) -> List[ChatCompletionMessageToolCall]:
|
||||
|
|
@ -379,13 +419,72 @@ class ToolPermissionGuardrail(CustomGuardrail):
|
|||
"""
|
||||
tool_calls = []
|
||||
|
||||
for choice in response.choices:
|
||||
for choice_index, choice in enumerate(response.choices):
|
||||
if isinstance(choice, Choices):
|
||||
for tool in choice.message.tool_calls or []:
|
||||
tool_calls.append(tool)
|
||||
legacy_tool_call = self._legacy_function_call_to_tool_call(
|
||||
getattr(choice.message, "function_call", None), choice_index
|
||||
)
|
||||
if legacy_tool_call is not None:
|
||||
tool_calls.append(legacy_tool_call)
|
||||
|
||||
return tool_calls
|
||||
|
||||
def _get_request_tool_name(self, tool: Any) -> tuple[Optional[str], Optional[str]]:
|
||||
tool_type = self._get_mapping_value(tool, "type")
|
||||
if tool_type != "function":
|
||||
return None, tool_type
|
||||
|
||||
function = self._get_mapping_value(tool, "function")
|
||||
tool_name = self._get_mapping_value(function, "name")
|
||||
return tool_name, tool_type
|
||||
|
||||
def _get_legacy_function_name(self, function: Any) -> Optional[str]:
|
||||
return self._get_mapping_value(function, "name")
|
||||
|
||||
def _get_named_tool_choice(self, data: dict) -> Optional[str]:
|
||||
tool_choice = data.get("tool_choice")
|
||||
if not tool_choice or tool_choice in ("auto", "none", "required"):
|
||||
return None
|
||||
if isinstance(tool_choice, str):
|
||||
return tool_choice
|
||||
if self._get_mapping_value(tool_choice, "type") != "function":
|
||||
return None
|
||||
return self._get_mapping_value(
|
||||
self._get_mapping_value(tool_choice, "function"), "name"
|
||||
)
|
||||
|
||||
def _get_named_function_call(self, data: dict) -> Optional[str]:
|
||||
function_call = data.get("function_call")
|
||||
if not function_call or function_call in ("auto", "none"):
|
||||
return None
|
||||
if isinstance(function_call, str):
|
||||
return function_call
|
||||
return self._get_mapping_value(function_call, "name")
|
||||
|
||||
def _collect_request_tools(self, data: dict) -> List[tuple[str, Optional[str]]]:
|
||||
request_tools: List[tuple[str, Optional[str]]] = []
|
||||
|
||||
for tool in data.get("tools") or []:
|
||||
tool_name, tool_type = self._get_request_tool_name(tool)
|
||||
if tool_name is not None:
|
||||
request_tools.append((tool_name, tool_type))
|
||||
|
||||
for function in data.get("functions") or []:
|
||||
function_name = self._get_legacy_function_name(function)
|
||||
if function_name is not None:
|
||||
request_tools.append((function_name, "function"))
|
||||
|
||||
for forced_tool_name in (
|
||||
self._get_named_tool_choice(data),
|
||||
self._get_named_function_call(data),
|
||||
):
|
||||
if forced_tool_name is not None:
|
||||
request_tools.append((forced_tool_name, "function"))
|
||||
|
||||
return request_tools
|
||||
|
||||
def _modify_request_with_permission_errors(
|
||||
self,
|
||||
data: dict,
|
||||
|
|
@ -410,19 +509,32 @@ class ToolPermissionGuardrail(CustomGuardrail):
|
|||
for tool_use in denied_tool_names:
|
||||
error_tool_names.add(tool_use)
|
||||
|
||||
# Modify the tools
|
||||
tools: Optional[List[ChatCompletionToolParam]] = data.get("tools")
|
||||
if tools is None:
|
||||
return data
|
||||
|
||||
new_tools = []
|
||||
for tool in tools:
|
||||
if tool["type"] != "function":
|
||||
continue
|
||||
tool_name: str = tool["function"]["name"]
|
||||
if tool_name not in error_tool_names:
|
||||
if tools is not None:
|
||||
new_tools = []
|
||||
for tool in tools:
|
||||
tool_name, tool_type = self._get_request_tool_name(tool)
|
||||
if tool_type == "function" and tool_name in error_tool_names:
|
||||
continue
|
||||
new_tools.append(tool)
|
||||
data["tools"] = new_tools
|
||||
data["tools"] = new_tools
|
||||
|
||||
functions = data.get("functions")
|
||||
if functions is not None:
|
||||
data["functions"] = [
|
||||
function
|
||||
for function in functions
|
||||
if self._get_legacy_function_name(function) not in error_tool_names
|
||||
]
|
||||
|
||||
named_tool_choice = self._get_named_tool_choice(data)
|
||||
if named_tool_choice in error_tool_names:
|
||||
data["tool_choice"] = "none"
|
||||
|
||||
named_function_call = self._get_named_function_call(data)
|
||||
if named_function_call in error_tool_names:
|
||||
data["function_call"] = "none"
|
||||
|
||||
return data
|
||||
|
||||
def _create_permission_error_result(
|
||||
|
|
@ -472,7 +584,7 @@ class ToolPermissionGuardrail(CustomGuardrail):
|
|||
error_results[tool_use.id] = error_result
|
||||
|
||||
# Modify the response content
|
||||
for choice in response.choices:
|
||||
for choice_index, choice in enumerate(response.choices):
|
||||
if isinstance(choice, Choices):
|
||||
filtered_tool_calls = []
|
||||
error_messages = []
|
||||
|
|
@ -490,6 +602,15 @@ class ToolPermissionGuardrail(CustomGuardrail):
|
|||
filtered_tool_calls if filtered_tool_calls else None
|
||||
)
|
||||
|
||||
legacy_tool_call = self._legacy_function_call_to_tool_call(
|
||||
getattr(choice.message, "function_call", None), choice_index
|
||||
)
|
||||
if legacy_tool_call is not None:
|
||||
legacy_error_result = error_results.get(legacy_tool_call.id)
|
||||
if legacy_error_result is not None:
|
||||
choice.message.function_call = None
|
||||
error_messages.append(legacy_error_result.content)
|
||||
|
||||
# Add error messages to content
|
||||
if error_messages:
|
||||
existing_content = choice.message.content
|
||||
|
|
@ -519,21 +640,16 @@ class ToolPermissionGuardrail(CustomGuardrail):
|
|||
if self.should_run_guardrail(data=data, event_type=event_type) is not True:
|
||||
return data
|
||||
|
||||
new_tools: Optional[List[ChatCompletionToolParam]] = data.get("tools")
|
||||
if new_tools is None:
|
||||
new_tools = self._collect_request_tools(data)
|
||||
if not new_tools:
|
||||
verbose_proxy_logger.warning(
|
||||
"Tool Permission Guardrail: not running guardrail. No tools in data"
|
||||
"Tool Permission Guardrail: not running guardrail. No tools or functions in data"
|
||||
)
|
||||
return data
|
||||
|
||||
# Check permissions for each tool
|
||||
denied_tool_names = []
|
||||
for tool in new_tools:
|
||||
if tool["type"] != "function":
|
||||
continue
|
||||
tool_name: str = tool["function"]["name"]
|
||||
tool_type: Optional[str] = tool.get("type")
|
||||
|
||||
for tool_name, tool_type in new_tools:
|
||||
is_allowed, _, message = self._check_tool_permission(tool_name, tool_type)
|
||||
|
||||
if not is_allowed and message is not None:
|
||||
|
|
|
|||
|
|
@ -29,6 +29,10 @@ ILLEGAL_DISPLAY_PARAMS = [
|
|||
"exception", # internal; not JSON-serializable, never for display
|
||||
"litellm_metadata", # internal tracking metadata with auth objects; not for display
|
||||
]
|
||||
# Provider routing fields. Allowed for proxy admins so they can see which
|
||||
# region/version a deployment is checking; gated at the endpoint layer for
|
||||
# non-admin callers (see _strip_admin_only_fields_from_health_result).
|
||||
ADMIN_ONLY_HEALTH_DISPLAY_PARAMS = ("api_base", "api_version")
|
||||
|
||||
MINIMAL_DISPLAY_PARAMS = ["model", "mode_error"]
|
||||
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ from litellm.proxy._types import (
|
|||
CallInfo,
|
||||
EnterpriseLicenseData,
|
||||
Litellm_EntityType,
|
||||
LitellmUserRoles,
|
||||
ProxyErrorTypes,
|
||||
ProxyException,
|
||||
UserAPIKeyAuth,
|
||||
|
|
@ -28,6 +29,7 @@ from litellm.proxy._types import (
|
|||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
|
||||
from litellm.proxy.health_check import (
|
||||
ADMIN_ONLY_HEALTH_DISPLAY_PARAMS,
|
||||
_clean_endpoint_data,
|
||||
_update_litellm_params_for_health_check,
|
||||
perform_health_check,
|
||||
|
|
@ -723,6 +725,90 @@ async def _save_background_health_checks_to_db(
|
|||
# Continue execution - don't let database save failure break health checks
|
||||
|
||||
|
||||
_PROXY_ADMIN_ROLES = frozenset(
|
||||
{
|
||||
LitellmUserRoles.PROXY_ADMIN.value,
|
||||
# View-only admins are operators (oncall, support); they need the
|
||||
# routing fields (api_base, api_version) to diagnose health and tell
|
||||
# which provider region a check is hitting. They cannot mutate config
|
||||
# so granting them the read-only view is safe.
|
||||
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _is_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> bool:
|
||||
"""
|
||||
Return True if the caller has a proxy-admin role (full or view-only).
|
||||
|
||||
user_role on UserAPIKeyAuth can be either a LitellmUserRoles enum or its
|
||||
string value depending on how the auth path constructed the object, so we
|
||||
compare against the raw value rather than the enum identity.
|
||||
"""
|
||||
role = user_api_key_dict.user_role
|
||||
if role is None:
|
||||
return False
|
||||
role_value = role.value if hasattr(role, "value") else role
|
||||
return role_value in _PROXY_ADMIN_ROLES
|
||||
|
||||
|
||||
def _strip_admin_only_fields_from_health_result(result: dict) -> dict:
|
||||
"""
|
||||
Return a copy of the /health response with provider routing fields
|
||||
(``api_base``, ``api_version``) removed from each healthy/unhealthy
|
||||
endpoint entry. Used to hide those fields from non-admin callers while
|
||||
still showing them which deployments they own and whether each one is
|
||||
healthy. Proxy admins receive the unmodified result.
|
||||
"""
|
||||
out = dict(result)
|
||||
drop = set(ADMIN_ONLY_HEALTH_DISPLAY_PARAMS)
|
||||
for key in ("healthy_endpoints", "unhealthy_endpoints"):
|
||||
eps = out.get(key)
|
||||
if isinstance(eps, list):
|
||||
out[key] = [
|
||||
(
|
||||
{k: v for k, v in ep.items() if k not in drop}
|
||||
if isinstance(ep, dict)
|
||||
else ep
|
||||
)
|
||||
for ep in eps
|
||||
]
|
||||
return out
|
||||
|
||||
|
||||
def _filter_health_check_results_by_model_ids(
|
||||
results: dict, allowed_model_ids: set
|
||||
) -> dict:
|
||||
"""
|
||||
Restrict a cached background health-check result dict to endpoints whose
|
||||
model_id is in ``allowed_model_ids``.
|
||||
|
||||
Endpoints without a model_id (e.g. CLI-model entries that predate the
|
||||
model_id wiring) are dropped conservatively — we cannot prove they belong
|
||||
to the caller, so they are excluded rather than leaked.
|
||||
|
||||
Each retained endpoint is shallow-copied before being returned, so any
|
||||
downstream transform (e.g. _strip_admin_only_fields_from_health_result)
|
||||
cannot accidentally mutate the shared ``health_check_results`` cache.
|
||||
"""
|
||||
healthy = [
|
||||
dict(ep)
|
||||
for ep in (results.get("healthy_endpoints") or [])
|
||||
if ep.get("model_id") in allowed_model_ids
|
||||
]
|
||||
unhealthy = [
|
||||
dict(ep)
|
||||
for ep in (results.get("unhealthy_endpoints") or [])
|
||||
if ep.get("model_id") in allowed_model_ids
|
||||
]
|
||||
return {
|
||||
"healthy_endpoints": healthy,
|
||||
"unhealthy_endpoints": unhealthy,
|
||||
"healthy_count": len(healthy),
|
||||
"unhealthy_count": len(unhealthy),
|
||||
}
|
||||
|
||||
|
||||
async def _perform_health_check_and_save(
|
||||
model_list,
|
||||
target_model,
|
||||
|
|
@ -771,6 +857,7 @@ async def _perform_health_check_and_save(
|
|||
|
||||
@router.get("/health", tags=["health"], dependencies=[Depends(user_api_key_auth)])
|
||||
async def health_endpoint(
|
||||
response: Response,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
model: Optional[str] = fastapi.Query(
|
||||
None, description="Specify the model name (optional)"
|
||||
|
|
@ -838,11 +925,26 @@ async def health_endpoint(
|
|||
detail={"error": f"Model with ID {model_id} not found"},
|
||||
)
|
||||
|
||||
is_admin = _is_proxy_admin(user_api_key_dict)
|
||||
|
||||
def _post_process(result: dict) -> dict:
|
||||
# api_base / api_version reveal which provider/region/internal host the
|
||||
# deployment talks to; only proxy admins receive them. Non-admin keys
|
||||
# still see model/model_id and the healthy/unhealthy status. We also
|
||||
# set a header so non-admin clients that previously parsed those
|
||||
# fields can detect the change programmatically.
|
||||
if is_admin:
|
||||
return result
|
||||
response.headers["Litellm-Health-Field-Notice"] = (
|
||||
"api_base and api_version are admin-only on this endpoint"
|
||||
)
|
||||
return _strip_admin_only_fields_from_health_result(result)
|
||||
|
||||
try:
|
||||
if llm_model_list is None:
|
||||
# if no router set, check if user set a model using litellm --model ollama/llama2
|
||||
if user_model is not None:
|
||||
return await _perform_health_check_and_save(
|
||||
cli_result = await _perform_health_check_and_save(
|
||||
model_list=[],
|
||||
target_model=None,
|
||||
cli_model=user_model,
|
||||
|
|
@ -853,20 +955,59 @@ async def health_endpoint(
|
|||
model_id=None, # CLI model doesn't have model_id
|
||||
max_concurrency=health_check_concurrency,
|
||||
)
|
||||
return _post_process(cli_result)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail={"error": "Model list not initialized"},
|
||||
)
|
||||
_llm_model_list = copy.deepcopy(llm_model_list)
|
||||
### FILTER MODELS FOR ONLY THOSE USER HAS ACCESS TO ###
|
||||
# Live path: scope by model_name (every deployment has one).
|
||||
# Cache path: scope by model_id (the cache is keyed on model_id).
|
||||
# Consequence: a deployment whose model_name the caller can access
|
||||
# but which lacks model_info.id will appear in the live /health
|
||||
# response but NOT in the background-cache /health response. This is
|
||||
# surfaced via the "warnings" field below so operators can fix the
|
||||
# missing model_info.id rather than guess at the discrepancy.
|
||||
if len(user_api_key_dict.models) > 0:
|
||||
pass
|
||||
else:
|
||||
pass #
|
||||
allowed_models = set(user_api_key_dict.models)
|
||||
_llm_model_list = [
|
||||
m for m in _llm_model_list if m.get("model_name") in allowed_models
|
||||
]
|
||||
if use_background_health_checks:
|
||||
return health_check_results
|
||||
if len(user_api_key_dict.models) > 0:
|
||||
allowed_model_ids = {
|
||||
(m.get("model_info") or {}).get("id")
|
||||
for m in _llm_model_list
|
||||
if (m.get("model_info") or {}).get("id")
|
||||
}
|
||||
filtered = _filter_health_check_results_by_model_ids(
|
||||
health_check_results, allowed_model_ids
|
||||
)
|
||||
if not allowed_model_ids:
|
||||
# Caller has accessible model_names but none of the
|
||||
# matching deployments expose a model_info.id, so the
|
||||
# cache filter (which keys on model_id) drops every
|
||||
# entry. Surface this both as a warning log and a
|
||||
# structured "warnings" field on the response so the
|
||||
# caller can distinguish "no deployments found" from
|
||||
# "deployments excluded due to missing model_info.id".
|
||||
verbose_proxy_logger.warning(
|
||||
"health_endpoint: scoped key %s has accessible models %s "
|
||||
"but none of the matching deployments carry a model_info.id; "
|
||||
"background health-check cache will return an empty result.",
|
||||
user_api_key_dict.user_id,
|
||||
list(user_api_key_dict.models),
|
||||
)
|
||||
filtered["warnings"] = [
|
||||
"Some accessible deployments are missing model_info.id "
|
||||
"and were excluded from this response. Ask a proxy admin "
|
||||
"to populate model_info.id for these models."
|
||||
]
|
||||
return _post_process(filtered)
|
||||
return _post_process(health_check_results)
|
||||
else:
|
||||
return await _perform_health_check_and_save(
|
||||
router_result = await _perform_health_check_and_save(
|
||||
model_list=_llm_model_list,
|
||||
target_model=target_model,
|
||||
cli_model=None,
|
||||
|
|
@ -877,6 +1018,7 @@ async def health_endpoint(
|
|||
model_id=model_id,
|
||||
max_concurrency=health_check_concurrency,
|
||||
)
|
||||
return _post_process(router_result)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(
|
||||
"litellm.proxy.proxy_server.py::health_endpoint(): Exception occured - {}".format(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import asyncio
|
||||
import json
|
||||
import time
|
||||
import urllib.parse
|
||||
from datetime import datetime
|
||||
from typing import Literal, Optional
|
||||
from urllib.parse import urlparse
|
||||
|
|
@ -203,8 +204,16 @@ class AssemblyAIPassthroughLoggingHandler:
|
|||
)
|
||||
if _api_key is None:
|
||||
raise ValueError("AssemblyAI API key not found")
|
||||
if (
|
||||
any(c in transcript_id for c in ("/", "\\", "#", "?"))
|
||||
or ".." in transcript_id
|
||||
):
|
||||
raise ValueError(
|
||||
f"Invalid transcript_id {transcript_id!r}: contains disallowed characters"
|
||||
)
|
||||
safe_transcript_id = urllib.parse.quote(transcript_id, safe="")
|
||||
try:
|
||||
url = f"{_base_url}/v2/transcript/{transcript_id}"
|
||||
url = f"{_base_url}/v2/transcript/{safe_transcript_id}"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {_api_key}",
|
||||
"Content-Type": "application/json",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
@ -13167,9 +13250,12 @@ async def update_config( # noqa: PLR0915
|
|||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
For Admin UI - allows admin to update config via UI
|
||||
For Admin UI - allows admin to update config via UI.
|
||||
|
||||
Currently supports modifying General Settings + LiteLLM settings
|
||||
Writes only the sections present in the request body to LiteLLM_Config rows
|
||||
(one row per top-level section). Sections the caller did not send are left
|
||||
untouched — this endpoint never persists pre-existing YAML values to DB as
|
||||
a side effect of an unrelated update.
|
||||
"""
|
||||
global llm_router, llm_model_list, general_settings, proxy_config, proxy_logging_obj, master_key, prisma_client
|
||||
try:
|
||||
|
|
@ -13177,109 +13263,96 @@ async def update_config( # noqa: PLR0915
|
|||
raise HTTPException(
|
||||
status_code=403, detail="Only proxy admins can update config"
|
||||
)
|
||||
import base64
|
||||
|
||||
"""
|
||||
- Update the ConfigTable DB
|
||||
- Run 'add_deployment'
|
||||
"""
|
||||
if prisma_client is None:
|
||||
raise Exception("No DB Connected")
|
||||
|
||||
if store_model_in_db is not True:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={
|
||||
"error": "Set `'STORE_MODEL_IN_DB='True'` in your env to enable this feature."
|
||||
async def _read_section(param_name: str) -> dict:
|
||||
row = await prisma_client.db.litellm_config.find_first(
|
||||
where={"param_name": param_name}
|
||||
)
|
||||
if row is None or row.param_value is None:
|
||||
return {}
|
||||
return dict(row.param_value)
|
||||
|
||||
async def _upsert_section(param_name: str, value: dict) -> None:
|
||||
serialized = json.dumps(value)
|
||||
await prisma_client.db.litellm_config.upsert(
|
||||
where={"param_name": param_name},
|
||||
data={
|
||||
"create": {"param_name": param_name, "param_value": serialized},
|
||||
"update": {"param_value": serialized},
|
||||
},
|
||||
)
|
||||
# invalidate the DualCache entry so the next reader (this process
|
||||
# or any other proxy in the cluster) goes to DB.
|
||||
await invalidate_config_param(param_name)
|
||||
|
||||
updated_settings = config_info.json(exclude_none=True)
|
||||
updated_settings = prisma_client.jsonify_object(updated_settings)
|
||||
for k, v in updated_settings.items():
|
||||
if k == "router_settings":
|
||||
await prisma_client.db.litellm_config.upsert(
|
||||
where={"param_name": k},
|
||||
data={
|
||||
"create": {"param_name": k, "param_value": v},
|
||||
"update": {"param_value": v},
|
||||
},
|
||||
)
|
||||
await invalidate_config_param(k)
|
||||
|
||||
### OLD LOGIC [TODO] MOVE TO DB ###
|
||||
|
||||
# Load existing config
|
||||
config = await proxy_config.get_config()
|
||||
verbose_proxy_logger.debug("Loaded config: %s", config)
|
||||
|
||||
# update the general settings
|
||||
# general_settings: merge per-key, with the alert_to_webhook_url side
|
||||
# effect of auto-enabling slack alerting.
|
||||
if config_info.general_settings is not None:
|
||||
config.setdefault("general_settings", {})
|
||||
updated_general_settings = config_info.general_settings.dict(
|
||||
exclude_none=True
|
||||
)
|
||||
|
||||
_existing_settings = config["general_settings"]
|
||||
for k, v in updated_general_settings.items():
|
||||
# overwrite existing settings with updated values
|
||||
existing = await _read_section("general_settings")
|
||||
updates = config_info.general_settings.dict(exclude_none=True)
|
||||
for k, v in updates.items():
|
||||
if k == "alert_to_webhook_url":
|
||||
# check if slack is already enabled. if not, enable it
|
||||
if "alerting" not in _existing_settings:
|
||||
_existing_settings = {"alerting": ["slack"]}
|
||||
elif isinstance(_existing_settings["alerting"], list):
|
||||
if "slack" not in _existing_settings["alerting"]:
|
||||
_existing_settings["alerting"].append("slack")
|
||||
_existing_settings[k] = v
|
||||
config["general_settings"] = _existing_settings
|
||||
if "alerting" not in existing:
|
||||
existing["alerting"] = ["slack"]
|
||||
elif (
|
||||
isinstance(existing["alerting"], list)
|
||||
and "slack" not in existing["alerting"]
|
||||
):
|
||||
existing["alerting"].append("slack")
|
||||
existing[k] = v
|
||||
await _upsert_section("general_settings", existing)
|
||||
|
||||
# environment_variables: encrypt request values, then merge into existing.
|
||||
if config_info.environment_variables is not None:
|
||||
config.setdefault("environment_variables", {})
|
||||
_updated_environment_variables = config_info.environment_variables
|
||||
existing = await _read_section("environment_variables")
|
||||
for k, v in config_info.environment_variables.items():
|
||||
existing[k] = encrypt_value_helper(value=v)
|
||||
await _upsert_section("environment_variables", existing)
|
||||
|
||||
# encrypt updated_environment_variables #
|
||||
for k, v in _updated_environment_variables.items():
|
||||
encrypted_value = encrypt_value_helper(value=v)
|
||||
_updated_environment_variables[k] = encrypted_value
|
||||
|
||||
_existing_env_variables = config["environment_variables"]
|
||||
|
||||
for k, v in _updated_environment_variables.items():
|
||||
# overwrite existing env variables with updated values
|
||||
_existing_env_variables[k] = _updated_environment_variables[k]
|
||||
|
||||
# update the litellm settings
|
||||
# litellm_settings: merge existing + request, request wins (matching
|
||||
# router_settings semantics — the caller's value for any given key is
|
||||
# what gets persisted). success_callback is special-cased: it is
|
||||
# always normalized + deduped, and unioned with any existing list,
|
||||
# because callbacks are additive (callers send the new entry, not
|
||||
# the full set). Normalizing on every write — not only when an
|
||||
# existing entry is present — keeps the DB free of mixed-case
|
||||
# entries that delete_callback (lowercase lookup) cannot find.
|
||||
if config_info.litellm_settings is not None:
|
||||
config.setdefault("litellm_settings", {})
|
||||
updated_litellm_settings = config_info.litellm_settings
|
||||
config["litellm_settings"] = {
|
||||
**updated_litellm_settings,
|
||||
**config["litellm_settings"],
|
||||
}
|
||||
existing = await _read_section("litellm_settings")
|
||||
updated_litellm_settings = dict(config_info.litellm_settings)
|
||||
|
||||
# if litellm.success_callback in updated_litellm_settings and config["litellm_settings"]
|
||||
if (
|
||||
"success_callback" in updated_litellm_settings
|
||||
and "success_callback" in config["litellm_settings"]
|
||||
):
|
||||
# check both success callback are lists
|
||||
if isinstance(
|
||||
config["litellm_settings"]["success_callback"], list
|
||||
) and isinstance(updated_litellm_settings["success_callback"], list):
|
||||
updated_success_callbacks_normalized = normalize_callback_names(
|
||||
updated_litellm_settings["success_callback"]
|
||||
)
|
||||
combined_success_callback = (
|
||||
config["litellm_settings"]["success_callback"]
|
||||
+ updated_success_callbacks_normalized
|
||||
)
|
||||
combined_success_callback = list(set(combined_success_callback))
|
||||
config["litellm_settings"][
|
||||
"success_callback"
|
||||
] = combined_success_callback
|
||||
incoming_cb = updated_litellm_settings.get("success_callback")
|
||||
if isinstance(incoming_cb, list):
|
||||
updated_litellm_settings["success_callback"] = normalize_callback_names(
|
||||
incoming_cb
|
||||
)
|
||||
|
||||
# Save the updated config
|
||||
await proxy_config.save_config(new_config=config)
|
||||
merged = {**existing, **updated_litellm_settings}
|
||||
|
||||
incoming_cb = updated_litellm_settings.get("success_callback")
|
||||
existing_cb = existing.get("success_callback")
|
||||
if isinstance(incoming_cb, list):
|
||||
if isinstance(existing_cb, list):
|
||||
# Normalize the existing list too — a row written by a
|
||||
# different code path may still hold mixed-case names,
|
||||
# which would otherwise dedup-miss against the lowercase
|
||||
# incoming entries.
|
||||
merged["success_callback"] = list(
|
||||
set(normalize_callback_names(existing_cb) + incoming_cb)
|
||||
)
|
||||
else:
|
||||
merged["success_callback"] = list(set(incoming_cb))
|
||||
|
||||
await _upsert_section("litellm_settings", merged)
|
||||
|
||||
# router_settings: merge existing + request, request wins.
|
||||
if config_info.router_settings is not None:
|
||||
existing = await _read_section("router_settings")
|
||||
updates = config_info.router_settings.dict(exclude_none=True)
|
||||
await _upsert_section("router_settings", {**existing, **updates})
|
||||
|
||||
await proxy_config.add_deployment(
|
||||
prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj
|
||||
|
|
|
|||
|
|
@ -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] = {}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -6,6 +6,13 @@ from typing_extensions import (
|
|||
TypedDict,
|
||||
)
|
||||
|
||||
from litellm.types.llms.openai import EmbeddingInput
|
||||
|
||||
# Gemini supports nested-list inputs (e.g. [["text", "image"]]) as an explicit
|
||||
# opt-in for combined embeddings — a provider-specific extension of the
|
||||
# OpenAI-faithful EmbeddingInput shape.
|
||||
GeminiEmbeddingInput = Union[EmbeddingInput, List[List[str]]]
|
||||
|
||||
|
||||
class FunctionResponse(TypedDict):
|
||||
name: str
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -34894,6 +34894,20 @@
|
|||
"supports_tool_choice": true,
|
||||
"source": "https://aws.amazon.com/bedrock/pricing/"
|
||||
},
|
||||
"zai.glm-5": {
|
||||
"input_cost_per_token": 1e-06,
|
||||
"litellm_provider": "bedrock_converse",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 3.2e-06,
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"source": "https://aws.amazon.com/bedrock/pricing/"
|
||||
},
|
||||
"zai.glm-4.7-flash": {
|
||||
"input_cost_per_token": 7e-08,
|
||||
"litellm_provider": "bedrock_converse",
|
||||
|
|
|
|||
|
|
@ -477,3 +477,4 @@ def test_get_llm_provider_use_proxy_arg_true_with_direct_args():
|
|||
assert provider == "litellm_proxy"
|
||||
assert key == arg_api_key # Should use the argument key
|
||||
assert base == arg_api_base # Should use the argument base
|
||||
|
||||
|
|
|
|||
|
|
@ -134,3 +134,62 @@ def test_is_assemblyai_route():
|
|||
== False
|
||||
)
|
||||
assert handler.is_assemblyai_route("") == False
|
||||
|
||||
|
||||
# --- Security: SSRF via transcript_id path traversal ---
|
||||
|
||||
|
||||
def test_get_assembly_transcript_rejects_slash_in_id(assembly_handler):
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
|
||||
return_value="test-key",
|
||||
):
|
||||
with pytest.raises(ValueError, match="disallowed characters"):
|
||||
assembly_handler._get_assembly_transcript("../../admin/credentials")
|
||||
|
||||
|
||||
def test_get_assembly_transcript_rejects_dotdot_in_id(assembly_handler):
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
|
||||
return_value="test-key",
|
||||
):
|
||||
with pytest.raises(ValueError, match="disallowed characters"):
|
||||
assembly_handler._get_assembly_transcript("..evil")
|
||||
|
||||
|
||||
def test_get_assembly_transcript_rejects_fragment_in_id(assembly_handler):
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
|
||||
return_value="test-key",
|
||||
):
|
||||
with pytest.raises(ValueError, match="disallowed characters"):
|
||||
assembly_handler._get_assembly_transcript("abc#suffix")
|
||||
|
||||
|
||||
def test_get_assembly_transcript_rejects_query_in_id(assembly_handler):
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
|
||||
return_value="test-key",
|
||||
):
|
||||
with pytest.raises(ValueError, match="disallowed characters"):
|
||||
assembly_handler._get_assembly_transcript("abc?x=1")
|
||||
|
||||
|
||||
def test_get_assembly_transcript_allows_valid_id(
|
||||
assembly_handler, mock_transcript_response
|
||||
):
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
|
||||
return_value="test-key",
|
||||
):
|
||||
with patch("httpx.get") as mock_get:
|
||||
mock_get.return_value.json.return_value = mock_transcript_response
|
||||
mock_get.return_value.raise_for_status.return_value = None
|
||||
|
||||
transcript = assembly_handler._get_assembly_transcript(
|
||||
"abc123-valid-id_xyz"
|
||||
)
|
||||
assert transcript == mock_transcript_response
|
||||
called_url = mock_get.call_args[0][0]
|
||||
assert "abc123-valid-id_xyz" in called_url
|
||||
assert ".." not in called_url
|
||||
|
|
|
|||
|
|
@ -2768,40 +2768,40 @@ async def test_update_config_success_callback_normalization():
|
|||
import litellm.proxy.proxy_server as proxy_server
|
||||
from litellm.proxy._types import ConfigYAML
|
||||
|
||||
# Ensure feature is enabled and prisma_client is set
|
||||
setattr(proxy_server, "store_model_in_db", True)
|
||||
setattr(proxy_server, "proxy_logging_obj", MagicMock())
|
||||
|
||||
existing_litellm_settings = {"success_callback": ["langfuse"]}
|
||||
|
||||
class FakeRow:
|
||||
def __init__(self, name, value):
|
||||
self.param_name = name
|
||||
self.param_value = value
|
||||
|
||||
upserted = {}
|
||||
|
||||
async def fake_find_first(where=None):
|
||||
if where and where.get("param_name") == "litellm_settings":
|
||||
return FakeRow("litellm_settings", existing_litellm_settings)
|
||||
return None
|
||||
|
||||
async def fake_upsert(where=None, data=None):
|
||||
upserted[where["param_name"]] = json.loads(data["update"]["param_value"])
|
||||
|
||||
class MockPrisma:
|
||||
def __init__(self):
|
||||
self.db = MagicMock()
|
||||
self.db.litellm_config = MagicMock()
|
||||
self.db.litellm_config.upsert = AsyncMock()
|
||||
|
||||
# proxy_server.update_config expects this to be sync returning a dict
|
||||
def jsonify_object(self, obj):
|
||||
return obj
|
||||
self.db.litellm_config.find_first = AsyncMock(side_effect=fake_find_first)
|
||||
self.db.litellm_config.upsert = AsyncMock(side_effect=fake_upsert)
|
||||
|
||||
setattr(proxy_server, "prisma_client", MockPrisma())
|
||||
|
||||
class MockProxyConfig:
|
||||
def __init__(self):
|
||||
self.saved_config = None
|
||||
|
||||
async def get_config(self):
|
||||
# Existing config has one lowercase callback already
|
||||
return {"litellm_settings": {"success_callback": ["langfuse"]}}
|
||||
|
||||
async def save_config(self, new_config: dict):
|
||||
self.saved_config = new_config
|
||||
|
||||
async def add_deployment(self, prisma_client=None, proxy_logging_obj=None):
|
||||
return None
|
||||
|
||||
mock_proxy_config = MockProxyConfig()
|
||||
setattr(proxy_server, "proxy_config", mock_proxy_config)
|
||||
setattr(proxy_server, "proxy_config", MockProxyConfig())
|
||||
|
||||
# Update config with mixed-case callbacks - expect normalization to lowercase
|
||||
config_update = ConfigYAML(litellm_settings={"success_callback": ["SQS", "sQs"]})
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
|
||||
|
|
@ -2810,9 +2810,10 @@ async def test_update_config_success_callback_normalization():
|
|||
)
|
||||
await proxy_server.update_config(config_update, user_api_key_dict=admin_user)
|
||||
|
||||
saved = mock_proxy_config.saved_config
|
||||
assert saved is not None, "save_config was not called"
|
||||
callbacks = saved["litellm_settings"]["success_callback"]
|
||||
assert (
|
||||
"litellm_settings" in upserted
|
||||
), "litellm_config.upsert was not called for litellm_settings"
|
||||
callbacks = upserted["litellm_settings"]["success_callback"]
|
||||
|
||||
# Deduped and normalized
|
||||
assert "sqs" in callbacks
|
||||
|
|
|
|||
|
|
@ -280,3 +280,55 @@ class TestDynamicProjectNameOnSpan:
|
|||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
|
||||
# --- Security: SSRF via prompt_version_id path traversal ---
|
||||
|
||||
|
||||
def test_arize_phoenix_client_sanitize_id_rejects_traversal():
|
||||
from litellm.integrations.arize.arize_phoenix_client import _sanitize_id
|
||||
|
||||
# dotdot without slashes
|
||||
with pytest.raises(ValueError, match="path traversal"):
|
||||
_sanitize_id("..something")
|
||||
# full traversal (slash caught first)
|
||||
with pytest.raises(ValueError, match="disallowed characters"):
|
||||
_sanitize_id("../../projects")
|
||||
|
||||
|
||||
def test_arize_phoenix_client_sanitize_id_rejects_slash():
|
||||
from litellm.integrations.arize.arize_phoenix_client import _sanitize_id
|
||||
|
||||
with pytest.raises(ValueError, match="disallowed characters"):
|
||||
_sanitize_id("valid/extra")
|
||||
|
||||
|
||||
def test_arize_phoenix_client_sanitize_id_rejects_fragment():
|
||||
from litellm.integrations.arize.arize_phoenix_client import _sanitize_id
|
||||
|
||||
with pytest.raises(ValueError, match="disallowed characters"):
|
||||
_sanitize_id("abc#suffix")
|
||||
|
||||
|
||||
def test_arize_phoenix_client_sanitize_id_rejects_query():
|
||||
from litellm.integrations.arize.arize_phoenix_client import _sanitize_id
|
||||
|
||||
with pytest.raises(ValueError, match="disallowed characters"):
|
||||
_sanitize_id("abc?x=1")
|
||||
|
||||
|
||||
def test_arize_phoenix_client_sanitize_id_allows_uuid():
|
||||
from litellm.integrations.arize.arize_phoenix_client import _sanitize_id
|
||||
|
||||
uid = "550e8400-e29b-41d4-a716-446655440000"
|
||||
assert _sanitize_id(uid) == uid
|
||||
|
||||
|
||||
def test_arize_phoenix_client_get_prompt_version_rejects_traversal():
|
||||
from litellm.integrations.arize.arize_phoenix_client import ArizePhoenixClient
|
||||
|
||||
client = ArizePhoenixClient(
|
||||
api_key="test-key", api_base="https://app.phoenix.arize.com"
|
||||
)
|
||||
with pytest.raises(ValueError, match="disallowed characters"):
|
||||
client.get_prompt_version("../../projects")
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ sys.path.insert(
|
|||
|
||||
import litellm
|
||||
from litellm.integrations.bitbucket import BitBucketPromptManager
|
||||
from litellm.integrations.bitbucket.bitbucket_client import _sanitize_file_path
|
||||
|
||||
|
||||
@patch("litellm.integrations.bitbucket.bitbucket_prompt_manager.BitBucketClient")
|
||||
|
|
@ -370,3 +371,45 @@ def test_bitbucket_prompt_manager_list_templates(mock_client_class):
|
|||
templates = manager.prompt_manager.list_templates()
|
||||
assert isinstance(templates, list)
|
||||
assert "test_prompt" in templates
|
||||
|
||||
|
||||
# --- Security: path traversal / SSRF ---
|
||||
|
||||
|
||||
def test_sanitize_file_path_rejects_traversal():
|
||||
with pytest.raises(ValueError, match="path traversal"):
|
||||
_sanitize_file_path("../../etc/passwd")
|
||||
|
||||
|
||||
def test_sanitize_file_path_rejects_fragment():
|
||||
with pytest.raises(ValueError, match="URL special characters"):
|
||||
_sanitize_file_path("secret#.prompt")
|
||||
|
||||
|
||||
def test_sanitize_file_path_rejects_query():
|
||||
with pytest.raises(ValueError, match="URL special characters"):
|
||||
_sanitize_file_path("secret?.prompt")
|
||||
|
||||
|
||||
def test_sanitize_file_path_encodes_special_chars():
|
||||
result = _sanitize_file_path("prompts/my prompt.prompt")
|
||||
assert result == "prompts/my%20prompt.prompt"
|
||||
|
||||
|
||||
def test_sanitize_file_path_allows_normal_paths():
|
||||
assert _sanitize_file_path("prompts/my-prompt") == "prompts/my-prompt"
|
||||
assert _sanitize_file_path("simple") == "simple"
|
||||
|
||||
|
||||
def test_bitbucket_client_rejects_traversal_in_get_file_content():
|
||||
from litellm.integrations.bitbucket.bitbucket_client import BitBucketClient
|
||||
|
||||
client = BitBucketClient(
|
||||
{
|
||||
"workspace": "ws",
|
||||
"repository": "repo",
|
||||
"access_token": "tok",
|
||||
}
|
||||
)
|
||||
with pytest.raises(ValueError, match="path traversal"):
|
||||
client.get_file_content("../../admin/credentials")
|
||||
|
|
|
|||
|
|
@ -878,6 +878,39 @@ def test_sync_streaming_bad_request_not_midstream(logging_obj: Logging):
|
|||
assert "invalid maxOutputTokens" in str(excinfo.value)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_streaming_read_timeout_triggers_midstream_fallback(
|
||||
logging_obj: Logging,
|
||||
):
|
||||
"""A mid-stream httpx.ReadTimeout must wrap into MidStreamFallbackError so
|
||||
the Router's FallbackStreamWrapper can switch to a fallback model.
|
||||
|
||||
Previously __anext__ caught httpx.TimeoutException and re-raised it raw,
|
||||
which bypassed _handle_stream_fallback_error and prevented stream_timeout
|
||||
from triggering fallbacks the way connection-phase timeout does.
|
||||
"""
|
||||
import httpx
|
||||
|
||||
from litellm.exceptions import MidStreamFallbackError
|
||||
|
||||
async def _raise_read_timeout(**kwargs):
|
||||
raise httpx.ReadTimeout("Timeout on reading data from socket")
|
||||
|
||||
response = CustomStreamWrapper(
|
||||
completion_stream=None,
|
||||
model="gpt-4",
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider="openai",
|
||||
make_call=_raise_read_timeout,
|
||||
)
|
||||
|
||||
with pytest.raises(MidStreamFallbackError) as excinfo:
|
||||
await response.__anext__()
|
||||
|
||||
assert excinfo.value.is_pre_first_chunk is True
|
||||
assert isinstance(excinfo.value.original_exception, Exception)
|
||||
|
||||
|
||||
def test_streaming_handler_with_created_time_propagation(
|
||||
initialized_custom_stream_wrapper: CustomStreamWrapper, logging_obj: Logging
|
||||
):
|
||||
|
|
|
|||
|
|
@ -394,3 +394,76 @@ class TestHostAllowlist:
|
|||
|
||||
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
|
||||
validate_url("http://internal.corp/")
|
||||
|
||||
|
||||
# ── assert_same_origin ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
from litellm.litellm_core_utils.url_utils import assert_same_origin
|
||||
|
||||
|
||||
def test_assert_same_origin_matches_scheme_host_port():
|
||||
"""A polling URL on the same scheme + host + port as the api_base
|
||||
passes — the upstream is trusted; the URL it returned points back at
|
||||
the same upstream."""
|
||||
assert_same_origin(
|
||||
"https://api.example.com/v1/operations/abc",
|
||||
"https://api.example.com/v1/generate",
|
||||
)
|
||||
|
||||
|
||||
def test_assert_same_origin_treats_default_ports_as_explicit():
|
||||
"""``https://x/`` and ``https://x:443/`` are the same origin."""
|
||||
assert_same_origin("https://api.example.com/poll", "https://api.example.com:443/")
|
||||
assert_same_origin("https://api.example.com:443/poll", "https://api.example.com/")
|
||||
assert_same_origin("http://api.example.com/poll", "http://api.example.com:80/")
|
||||
|
||||
|
||||
def test_assert_same_origin_rejects_different_host():
|
||||
with pytest.raises(SSRFError, match="host"):
|
||||
assert_same_origin(
|
||||
"https://attacker.example.com/poll",
|
||||
"https://api.example.com/generate",
|
||||
)
|
||||
|
||||
|
||||
def test_assert_same_origin_rejects_different_scheme():
|
||||
with pytest.raises(SSRFError, match="scheme"):
|
||||
assert_same_origin(
|
||||
"http://api.example.com/poll", "https://api.example.com/generate"
|
||||
)
|
||||
|
||||
|
||||
def test_assert_same_origin_rejects_different_port():
|
||||
with pytest.raises(SSRFError, match="port"):
|
||||
assert_same_origin(
|
||||
"https://api.example.com:8443/poll", "https://api.example.com/generate"
|
||||
)
|
||||
|
||||
|
||||
def test_assert_same_origin_rejects_non_http_scheme():
|
||||
"""``file://`` polling URLs are rejected outright — the upstream
|
||||
should never return a non-HTTP scheme."""
|
||||
with pytest.raises(SSRFError, match="scheme"):
|
||||
assert_same_origin("file:///etc/passwd", "https://api.example.com/")
|
||||
|
||||
|
||||
def test_assert_same_origin_case_insensitive_host():
|
||||
assert_same_origin(
|
||||
"https://API.example.com/poll", "https://api.example.com/generate"
|
||||
)
|
||||
|
||||
|
||||
def test_assert_same_origin_error_message_does_not_leak_hostnames():
|
||||
"""Greptile P2: in the SSRF threat model the caller is the attacker.
|
||||
The error message must not echo the operator's expected host or the
|
||||
attacker-supplied candidate host back to the caller — only identify
|
||||
*which* component mismatched."""
|
||||
with pytest.raises(SSRFError) as exc:
|
||||
assert_same_origin(
|
||||
"https://attacker.example.com:1234/poll",
|
||||
"https://api.internal-corp.example/generate",
|
||||
)
|
||||
detail = str(exc.value)
|
||||
assert "attacker.example.com" not in detail
|
||||
assert "api.internal-corp.example" not in detail
|
||||
|
|
|
|||
177
tests/test_litellm/llms/test_polling_url_origin_match.py
Normal file
177
tests/test_litellm/llms/test_polling_url_origin_match.py
Normal file
|
|
@ -0,0 +1,177 @@
|
|||
"""
|
||||
VERIA-51: polling URLs returned by upstream APIs (Azure DALL-E,
|
||||
Azure Document Intelligence, Black Forest Labs) used to be followed
|
||||
without origin validation. The handlers attached the operator's API
|
||||
key to the polling request, so an attacker who could influence the
|
||||
upstream response (or a compromised upstream) could redirect the proxy
|
||||
to send credentials anywhere.
|
||||
|
||||
These tests assert each handler now rejects polling URLs that don't
|
||||
share an origin with the original request URL.
|
||||
"""
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
|
||||
# Azure DALL-E sync + async paths route through ``assert_same_origin``
|
||||
# the same way as the cases below. The helper itself is unit-tested in
|
||||
# ``tests/test_litellm/litellm_core_utils/test_url_utils.py``; the
|
||||
# tests here exercise the wiring at sites with simpler signatures.
|
||||
|
||||
|
||||
# ── Azure Document Intelligence polling ───────────────────────────────────────
|
||||
|
||||
|
||||
def test_azure_di_sync_rejects_cross_origin_polling():
|
||||
from litellm.llms.azure_ai.ocr.document_intelligence.transformation import (
|
||||
AzureDocumentIntelligenceOCRConfig,
|
||||
)
|
||||
|
||||
config = AzureDocumentIntelligenceOCRConfig()
|
||||
|
||||
raw_response = MagicMock()
|
||||
raw_response.status_code = 202
|
||||
raw_response.headers = {
|
||||
"Operation-Location": "https://attacker.example.com/results/xyz",
|
||||
}
|
||||
raw_response.request = MagicMock()
|
||||
raw_response.request.url = (
|
||||
"https://eastus.cognitiveservices.azure.com/documentintelligence/.../analyze"
|
||||
)
|
||||
raw_response.request.headers = {"Ocp-Apim-Subscription-Key": "leak-me"}
|
||||
|
||||
with pytest.raises(ValueError, match="rejected polling URL"):
|
||||
config.transform_ocr_response(
|
||||
model="azure-doc-intel",
|
||||
raw_response=raw_response,
|
||||
logging_obj=MagicMock(),
|
||||
request_data={},
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
encoding=None,
|
||||
response={},
|
||||
)
|
||||
|
||||
|
||||
# ── Black Forest Labs polling ─────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_bfl_image_generation_sync_rejects_cross_origin_polling():
|
||||
from litellm.llms.black_forest_labs.image_generation.handler import (
|
||||
BlackForestLabsImageGeneration,
|
||||
)
|
||||
|
||||
handler = BlackForestLabsImageGeneration()
|
||||
|
||||
initial_response = MagicMock()
|
||||
initial_response.status_code = 200
|
||||
initial_response.json = MagicMock(
|
||||
return_value={"polling_url": "https://attacker.example.com/get_result"}
|
||||
)
|
||||
initial_response.request = MagicMock()
|
||||
initial_response.request.url = "https://api.bfl.ai/v1/flux-pro"
|
||||
|
||||
sync_client = MagicMock()
|
||||
sync_client.get = MagicMock()
|
||||
|
||||
with pytest.raises(Exception, match="Rejected polling URL"):
|
||||
handler._poll_for_result_sync(
|
||||
initial_response=initial_response,
|
||||
headers={"x-key": "secret"},
|
||||
sync_client=sync_client,
|
||||
)
|
||||
|
||||
sync_client.get.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bfl_image_generation_async_rejects_cross_origin_polling():
|
||||
from litellm.llms.black_forest_labs.image_generation.handler import (
|
||||
BlackForestLabsImageGeneration,
|
||||
)
|
||||
|
||||
handler = BlackForestLabsImageGeneration()
|
||||
|
||||
initial_response = MagicMock()
|
||||
initial_response.status_code = 200
|
||||
initial_response.json = MagicMock(
|
||||
return_value={"polling_url": "https://attacker.example.com/get_result"}
|
||||
)
|
||||
initial_response.request = MagicMock()
|
||||
initial_response.request.url = "https://api.bfl.ai/v1/flux-pro"
|
||||
|
||||
async_client = MagicMock()
|
||||
async_client.get = MagicMock()
|
||||
|
||||
with pytest.raises(Exception, match="Rejected polling URL"):
|
||||
await handler._poll_for_result_async(
|
||||
initial_response=initial_response,
|
||||
headers={"x-key": "secret"},
|
||||
async_client=async_client,
|
||||
)
|
||||
|
||||
async_client.get.assert_not_called()
|
||||
|
||||
|
||||
def test_bfl_image_edit_sync_rejects_cross_origin_polling():
|
||||
from litellm.llms.black_forest_labs.image_edit.handler import (
|
||||
BlackForestLabsImageEdit,
|
||||
)
|
||||
|
||||
handler = BlackForestLabsImageEdit()
|
||||
|
||||
initial_response = MagicMock()
|
||||
initial_response.status_code = 200
|
||||
initial_response.json = MagicMock(
|
||||
return_value={"polling_url": "https://attacker.example.com/get_result"}
|
||||
)
|
||||
initial_response.request = MagicMock()
|
||||
initial_response.request.url = "https://api.bfl.ai/v1/flux-pro/edit"
|
||||
|
||||
sync_client = MagicMock()
|
||||
sync_client.get = MagicMock()
|
||||
|
||||
with pytest.raises(Exception, match="Rejected polling URL"):
|
||||
handler._poll_for_result_sync(
|
||||
initial_response=initial_response,
|
||||
headers={"x-key": "secret"},
|
||||
sync_client=sync_client,
|
||||
)
|
||||
|
||||
sync_client.get.assert_not_called()
|
||||
|
||||
|
||||
def test_bfl_image_generation_same_origin_polling_passes():
|
||||
"""Sanity check: when the polling URL shares origin with the original
|
||||
request, the origin check passes and polling proceeds."""
|
||||
from litellm.llms.black_forest_labs.image_generation.handler import (
|
||||
BlackForestLabsImageGeneration,
|
||||
)
|
||||
|
||||
handler = BlackForestLabsImageGeneration()
|
||||
|
||||
initial_response = MagicMock()
|
||||
initial_response.status_code = 200
|
||||
initial_response.json = MagicMock(
|
||||
return_value={"polling_url": "https://api.bfl.ai/v1/get_result?id=abc"}
|
||||
)
|
||||
initial_response.request = MagicMock()
|
||||
initial_response.request.url = "https://api.bfl.ai/v1/flux-pro"
|
||||
|
||||
sync_client = MagicMock()
|
||||
poll_response = MagicMock()
|
||||
poll_response.status_code = 200
|
||||
poll_response.json = MagicMock(return_value={"status": "Ready"})
|
||||
sync_client.get = MagicMock(return_value=poll_response)
|
||||
|
||||
result = handler._poll_for_result_sync(
|
||||
initial_response=initial_response,
|
||||
headers={"x-key": "secret"},
|
||||
sync_client=sync_client,
|
||||
)
|
||||
|
||||
sync_client.get.assert_called_once()
|
||||
assert result is poll_response
|
||||
|
|
@ -0,0 +1,290 @@
|
|||
"""
|
||||
Tests for Gemini batchEmbedContents transformation logic.
|
||||
|
||||
Covers:
|
||||
- Text-only inputs (single and batch)
|
||||
- Multimodal inputs (data URIs, GCS URLs, file references)
|
||||
- Mixed text + multimodal inputs
|
||||
- Response processing with correct indices
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation import (
|
||||
_build_part_for_input,
|
||||
_is_multimodal_input,
|
||||
process_response,
|
||||
transform_openai_input_gemini_content,
|
||||
transform_openai_input_gemini_embed_content,
|
||||
)
|
||||
from litellm.types.llms.vertex_ai import VertexAIBatchEmbeddingsResponseObject
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
|
||||
|
||||
IMAGE_DATA_URI = "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAgAAAAIAQMAAAD+wSzIAAAABlBMVEX///+/v7+jQ3Y5AAAADklEQVQI12P4AIX8EAgALgAD/aNpbtEAAAAASUVORK5CYII"
|
||||
GCS_URL = "gs://my-bucket/image.png"
|
||||
|
||||
|
||||
class TestIsMultimodalInput:
|
||||
def test_text_only_string(self):
|
||||
assert _is_multimodal_input("hello world") is False
|
||||
|
||||
def test_text_only_list(self):
|
||||
assert _is_multimodal_input(["hello", "world"]) is False
|
||||
|
||||
def test_data_uri(self):
|
||||
assert _is_multimodal_input([IMAGE_DATA_URI]) is True
|
||||
|
||||
def test_gcs_url(self):
|
||||
assert _is_multimodal_input([GCS_URL]) is True
|
||||
|
||||
def test_file_reference(self):
|
||||
assert _is_multimodal_input(["files/abc123"]) is True
|
||||
|
||||
def test_mixed_text_and_image(self):
|
||||
assert _is_multimodal_input(["hello", IMAGE_DATA_URI]) is True
|
||||
|
||||
def test_nested_text_is_not_multimodal(self):
|
||||
"""Nested list with text is not multimodal."""
|
||||
assert _is_multimodal_input([["text_a", "text_b"]]) is False
|
||||
|
||||
def test_nested_list_with_image_is_multimodal(self):
|
||||
assert _is_multimodal_input([["a red shoe", IMAGE_DATA_URI]]) is True
|
||||
|
||||
|
||||
class TestBuildPartForInput:
|
||||
def test_text_input(self):
|
||||
part = _build_part_for_input("hello")
|
||||
assert part["text"] == "hello"
|
||||
assert part.get("inline_data") is None
|
||||
|
||||
def test_data_uri_input(self):
|
||||
part = _build_part_for_input(IMAGE_DATA_URI)
|
||||
assert part.get("text") is None
|
||||
assert part["inline_data"] is not None
|
||||
assert part["inline_data"]["mime_type"] == "image/png"
|
||||
|
||||
def test_gcs_url_input(self):
|
||||
part = _build_part_for_input(GCS_URL)
|
||||
assert part.get("text") is None
|
||||
assert part["file_data"] is not None
|
||||
assert part["file_data"]["mime_type"] == "image/png"
|
||||
assert part["file_data"]["file_uri"] == GCS_URL
|
||||
|
||||
def test_file_reference_resolved(self):
|
||||
resolved = {"files/abc": {"mime_type": "image/jpeg", "uri": "https://example.com/abc"}}
|
||||
part = _build_part_for_input("files/abc", resolved_files=resolved)
|
||||
assert part["file_data"] is not None
|
||||
assert part["file_data"]["mime_type"] == "image/jpeg"
|
||||
|
||||
def test_file_reference_unresolved_raises(self):
|
||||
with pytest.raises(ValueError, match="not resolved"):
|
||||
_build_part_for_input("files/abc")
|
||||
|
||||
|
||||
class TestTransformOpenaiInputGeminiContent:
|
||||
"""Test that transform_openai_input_gemini_content creates separate requests per input."""
|
||||
|
||||
def test_single_text(self):
|
||||
result = transform_openai_input_gemini_content(
|
||||
input="hello", model="gemini-embedding-2-preview", optional_params={}
|
||||
)
|
||||
assert len(result["requests"]) == 1
|
||||
assert result["requests"][0]["content"]["parts"][0]["text"] == "hello"
|
||||
|
||||
def test_multiple_texts(self):
|
||||
result = transform_openai_input_gemini_content(
|
||||
input=["hello", "world"], model="gemini-embedding-2-preview", optional_params={}
|
||||
)
|
||||
assert len(result["requests"]) == 2
|
||||
assert result["requests"][0]["content"]["parts"][0]["text"] == "hello"
|
||||
assert result["requests"][1]["content"]["parts"][0]["text"] == "world"
|
||||
|
||||
def test_multimodal_inputs_are_separate_requests(self):
|
||||
"""Key regression test for #24209: each input becomes its own request."""
|
||||
result = transform_openai_input_gemini_content(
|
||||
input=["The food was delicious", IMAGE_DATA_URI],
|
||||
model="gemini-embedding-2-preview",
|
||||
optional_params={},
|
||||
)
|
||||
assert len(result["requests"]) == 2
|
||||
# First request is text
|
||||
assert result["requests"][0]["content"]["parts"][0]["text"] == "The food was delicious"
|
||||
# Second request is image
|
||||
assert result["requests"][1]["content"]["parts"][0]["inline_data"] is not None
|
||||
|
||||
def test_dimensions_mapped_to_output_dimensionality(self):
|
||||
result = transform_openai_input_gemini_content(
|
||||
input="hello",
|
||||
model="gemini-embedding-2-preview",
|
||||
optional_params={"dimensions": 256},
|
||||
)
|
||||
assert result["requests"][0]["outputDimensionality"] == 256
|
||||
|
||||
def test_model_name_prefixed(self):
|
||||
result = transform_openai_input_gemini_content(
|
||||
input="hello", model="gemini-embedding-2-preview", optional_params={}
|
||||
)
|
||||
assert result["requests"][0]["model"] == "models/gemini-embedding-2-preview"
|
||||
|
||||
def test_gcs_url_input(self):
|
||||
result = transform_openai_input_gemini_content(
|
||||
input=[GCS_URL], model="gemini-embedding-2-preview", optional_params={}
|
||||
)
|
||||
assert len(result["requests"]) == 1
|
||||
assert result["requests"][0]["content"]["parts"][0]["file_data"] is not None
|
||||
|
||||
def test_mixed_text_image_gcs(self):
|
||||
result = transform_openai_input_gemini_content(
|
||||
input=["hello", IMAGE_DATA_URI, GCS_URL],
|
||||
model="gemini-embedding-2-preview",
|
||||
optional_params={},
|
||||
)
|
||||
assert len(result["requests"]) == 3
|
||||
|
||||
def test_nested_input_combined_embedding(self):
|
||||
"""Nested list produces one request with multiple parts (combined embedding)."""
|
||||
result = transform_openai_input_gemini_content(
|
||||
input=[["a red shoe", IMAGE_DATA_URI]],
|
||||
model="gemini-embedding-2-preview",
|
||||
optional_params={},
|
||||
)
|
||||
assert len(result["requests"]) == 1
|
||||
parts = result["requests"][0]["content"]["parts"]
|
||||
assert len(parts) == 2
|
||||
assert parts[0]["text"] == "a red shoe"
|
||||
assert parts[1]["inline_data"] is not None
|
||||
|
||||
def test_mixed_nested_and_flat(self):
|
||||
"""Mixed nested + flat produces correct number of requests."""
|
||||
result = transform_openai_input_gemini_content(
|
||||
input=[["text", IMAGE_DATA_URI], "standalone"],
|
||||
model="gemini-embedding-2-preview",
|
||||
optional_params={},
|
||||
)
|
||||
assert len(result["requests"]) == 2
|
||||
# First: combined (2 parts)
|
||||
assert len(result["requests"][0]["content"]["parts"]) == 2
|
||||
# Second: standalone (1 part)
|
||||
assert len(result["requests"][1]["content"]["parts"]) == 1
|
||||
assert result["requests"][1]["content"]["parts"][0]["text"] == "standalone"
|
||||
|
||||
|
||||
class TestTransformOpenaiInputGeminiEmbedContent:
|
||||
"""Test transform_openai_input_gemini_embed_content (vertex_ai / embedContent path)."""
|
||||
|
||||
def test_text_and_image_combined(self):
|
||||
result = transform_openai_input_gemini_embed_content(
|
||||
input=["hello", IMAGE_DATA_URI],
|
||||
model="gemini-embedding-2-preview",
|
||||
optional_params={},
|
||||
)
|
||||
assert "content" in result
|
||||
parts = result["content"]["parts"]
|
||||
assert len(parts) == 2
|
||||
assert parts[0]["text"] == "hello"
|
||||
assert parts[1]["inline_data"] is not None
|
||||
|
||||
def test_gcs_url(self):
|
||||
result = transform_openai_input_gemini_embed_content(
|
||||
input=[GCS_URL],
|
||||
model="gemini-embedding-2-preview",
|
||||
optional_params={},
|
||||
)
|
||||
parts = result["content"]["parts"]
|
||||
assert len(parts) == 1
|
||||
assert parts[0]["file_data"]["file_uri"] == GCS_URL
|
||||
|
||||
def test_dimensions_mapped(self):
|
||||
result = transform_openai_input_gemini_embed_content(
|
||||
input="hello",
|
||||
model="gemini-embedding-2-preview",
|
||||
optional_params={"dimensions": 256},
|
||||
)
|
||||
assert result["outputDimensionality"] == 256
|
||||
|
||||
|
||||
class TestProcessResponse:
|
||||
"""Test that process_response sets correct indices."""
|
||||
|
||||
def test_single_embedding_index(self):
|
||||
predictions: VertexAIBatchEmbeddingsResponseObject = {
|
||||
"embeddings": [{"values": [0.1, 0.2]}]
|
||||
}
|
||||
model_response = EmbeddingResponse()
|
||||
result = process_response(
|
||||
input="hello",
|
||||
model_response=model_response,
|
||||
model="gemini-embedding-2-preview",
|
||||
_predictions=predictions,
|
||||
)
|
||||
assert len(result.data) == 1
|
||||
assert result.data[0]["index"] == 0
|
||||
|
||||
def test_multiple_embeddings_have_correct_indices(self):
|
||||
"""Regression test: indices should be 0, 1, 2... not all 0."""
|
||||
predictions: VertexAIBatchEmbeddingsResponseObject = {
|
||||
"embeddings": [
|
||||
{"values": [0.1, 0.2]},
|
||||
{"values": [0.3, 0.4]},
|
||||
{"values": [0.5, 0.6]},
|
||||
]
|
||||
}
|
||||
model_response = EmbeddingResponse()
|
||||
result = process_response(
|
||||
input=["a", "b", "c"],
|
||||
model_response=model_response,
|
||||
model="gemini-embedding-2-preview",
|
||||
_predictions=predictions,
|
||||
)
|
||||
assert len(result.data) == 3
|
||||
assert result.data[0]["index"] == 0
|
||||
assert result.data[1]["index"] == 1
|
||||
assert result.data[2]["index"] == 2
|
||||
|
||||
def test_multimodal_mixed_input(self):
|
||||
"""process_response works with mixed text + multimodal inputs."""
|
||||
predictions: VertexAIBatchEmbeddingsResponseObject = {
|
||||
"embeddings": [{"values": [0.1, 0.2]}, {"values": [0.3, 0.4]}]
|
||||
}
|
||||
result = process_response(
|
||||
input=["hello", IMAGE_DATA_URI],
|
||||
model_response=EmbeddingResponse(),
|
||||
model="gemini-embedding-2-preview",
|
||||
_predictions=predictions,
|
||||
)
|
||||
assert len(result.data) == 2
|
||||
assert result.data[0]["index"] == 0
|
||||
assert result.data[1]["index"] == 1
|
||||
# Should count tokens only for the text element, not the image
|
||||
assert result.usage.prompt_tokens > 0
|
||||
|
||||
def test_nested_input_token_counting(self):
|
||||
"""Nested list: only plain-text sub-elements should be counted."""
|
||||
predictions: VertexAIBatchEmbeddingsResponseObject = {
|
||||
"embeddings": [{"values": [0.1, 0.2]}]
|
||||
}
|
||||
result = process_response(
|
||||
input=[["a red shoe", IMAGE_DATA_URI]],
|
||||
model_response=EmbeddingResponse(),
|
||||
model="gemini-embedding-2-preview",
|
||||
_predictions=predictions,
|
||||
)
|
||||
assert len(result.data) == 1
|
||||
assert result.usage.prompt_tokens > 0
|
||||
|
||||
def test_nested_empty_list_raises(self):
|
||||
with pytest.raises(ValueError, match="must not be empty"):
|
||||
transform_openai_input_gemini_content(
|
||||
input=[[]],
|
||||
model="gemini-embedding-2-preview",
|
||||
optional_params={},
|
||||
)
|
||||
|
||||
def test_nested_non_string_element_raises(self):
|
||||
with pytest.raises(ValueError, match="must be strings"):
|
||||
transform_openai_input_gemini_content(
|
||||
input=[[["doubly", "nested"]]],
|
||||
model="gemini-embedding-2-preview",
|
||||
optional_params={},
|
||||
)
|
||||
|
|
@ -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
|
||||
|
||||
|
|
@ -964,3 +1166,129 @@ class TestIsRequestBodySafeBlocksEndpointTargetingFields:
|
|||
)
|
||||
is True
|
||||
)
|
||||
|
||||
|
||||
# ── is_request_body_safe nested-config recursion (VERIA-6) ────────────────────
|
||||
|
||||
|
||||
class TestIsRequestBodySafeNestedConfig:
|
||||
"""The Milvus vector store transformer unpacks
|
||||
``litellm_embedding_config`` as ``**kwargs`` into ``litellm.embedding(...)``
|
||||
— same SSRF / credential-exfil surface as a top-level ``api_base`` in
|
||||
the request body. ``is_request_body_safe`` must recurse into this
|
||||
nested dict so a banned param can't be smuggled in via nesting."""
|
||||
|
||||
def test_root_level_api_base_blocked_when_no_opt_in(self):
|
||||
"""Sanity check: pre-existing root-level enforcement still works."""
|
||||
with pytest.raises(ValueError, match="api_base"):
|
||||
is_request_body_safe(
|
||||
request_body={"api_base": "https://attacker.example.com"},
|
||||
general_settings={},
|
||||
llm_router=None,
|
||||
model="gpt-4",
|
||||
)
|
||||
|
||||
def test_nested_api_base_in_embedding_config_blocked(self):
|
||||
"""Smuggling ``api_base`` inside ``litellm_embedding_config`` is
|
||||
the VERIA-6 bypass — must be blocked by the recursive check."""
|
||||
with pytest.raises(ValueError, match="api_base"):
|
||||
is_request_body_safe(
|
||||
request_body={
|
||||
"litellm_embedding_config": {
|
||||
"api_base": "https://attacker.example.com",
|
||||
"api_key": "leaked-key",
|
||||
}
|
||||
},
|
||||
general_settings={},
|
||||
llm_router=None,
|
||||
model="milvus-store",
|
||||
)
|
||||
|
||||
def test_nested_langfuse_host_in_embedding_config_blocked(self):
|
||||
"""The recursion uses the *full* banned-param list, not a special
|
||||
subset — so any flag that's banned at the root is also banned
|
||||
when nested."""
|
||||
with pytest.raises(ValueError, match="langfuse_host"):
|
||||
is_request_body_safe(
|
||||
request_body={
|
||||
"litellm_embedding_config": {
|
||||
"langfuse_host": "https://attacker.example.com"
|
||||
}
|
||||
},
|
||||
general_settings={},
|
||||
llm_router=None,
|
||||
model="milvus-store",
|
||||
)
|
||||
|
||||
def test_nested_api_base_allowed_when_admin_opts_in(self):
|
||||
"""Admins who explicitly enable client-side credential passthrough
|
||||
keep the existing escape hatch — same UX as for root-level."""
|
||||
assert (
|
||||
is_request_body_safe(
|
||||
request_body={
|
||||
"litellm_embedding_config": {
|
||||
"api_base": "https://my-azure.example.com"
|
||||
}
|
||||
},
|
||||
general_settings={"allow_client_side_credentials": True},
|
||||
llm_router=None,
|
||||
model="milvus-store",
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
def test_safe_nested_config_accepted(self):
|
||||
"""A nested config without any banned params passes — there's no
|
||||
false-positive on legitimate ``api_version`` / model params."""
|
||||
assert (
|
||||
is_request_body_safe(
|
||||
request_body={
|
||||
"litellm_embedding_config": {
|
||||
"api_version": "2024-02-15-preview",
|
||||
}
|
||||
},
|
||||
general_settings={},
|
||||
llm_router=None,
|
||||
model="milvus-store",
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
def test_non_dict_nested_config_does_not_break_check(self):
|
||||
"""A bogus type for ``litellm_embedding_config`` (string, list,
|
||||
None) must not crash the validator — it should just fall through."""
|
||||
assert (
|
||||
is_request_body_safe(
|
||||
request_body={"litellm_embedding_config": "not-a-dict"},
|
||||
general_settings={},
|
||||
llm_router=None,
|
||||
model="x",
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
def test_deeply_nested_config_does_not_recurse(self):
|
||||
"""Greptile P1: ``is_request_body_safe`` is iterative single-level —
|
||||
a deeply-nested ``litellm_embedding_config`` cannot exhaust the
|
||||
Python call stack to trigger a 500 ``RecursionError``. Build a
|
||||
body 1000 levels deep; the validator must complete in O(1)
|
||||
descent."""
|
||||
body = {"litellm_embedding_config": {}}
|
||||
cur = body["litellm_embedding_config"]
|
||||
for _ in range(1000):
|
||||
cur["litellm_embedding_config"] = {}
|
||||
cur = cur["litellm_embedding_config"]
|
||||
# Banned param at the deepest level shouldn't be reached — single
|
||||
# level only.
|
||||
cur["api_base"] = "https://attacker.example.com"
|
||||
|
||||
# No exception raised: deeper levels aren't checked.
|
||||
assert (
|
||||
is_request_body_safe(
|
||||
request_body=body,
|
||||
general_settings={},
|
||||
llm_router=None,
|
||||
model="x",
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import json
|
||||
import os
|
||||
import sys
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock, patch
|
||||
|
||||
sys.path.insert(
|
||||
|
|
@ -34,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"
|
||||
|
|
@ -177,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"])
|
||||
|
|
|
|||
|
|
@ -231,6 +231,50 @@ class TestTokenUtilities:
|
|||
result = get_stored_api_key()
|
||||
assert result is None
|
||||
|
||||
def test_get_stored_api_key_base_url_match(self):
|
||||
"""Stored key is returned when expected_base_url matches stored origin"""
|
||||
token_data = {"key": "sk-prod", "base_url": "https://real-proxy.com"}
|
||||
with patch(
|
||||
"litellm.litellm_core_utils.cli_token_utils.load_cli_token",
|
||||
return_value=token_data,
|
||||
):
|
||||
assert (
|
||||
get_stored_api_key(expected_base_url="https://real-proxy.com")
|
||||
== "sk-prod"
|
||||
)
|
||||
|
||||
def test_get_stored_api_key_base_url_match_trailing_slash(self):
|
||||
"""Trailing slash on expected_base_url is normalised before comparison"""
|
||||
token_data = {"key": "sk-prod", "base_url": "https://real-proxy.com"}
|
||||
with patch(
|
||||
"litellm.litellm_core_utils.cli_token_utils.load_cli_token",
|
||||
return_value=token_data,
|
||||
):
|
||||
assert (
|
||||
get_stored_api_key(expected_base_url="https://real-proxy.com/")
|
||||
== "sk-prod"
|
||||
)
|
||||
|
||||
def test_get_stored_api_key_base_url_mismatch(self):
|
||||
"""Stored key is NOT returned when expected_base_url differs from stored origin"""
|
||||
token_data = {"key": "sk-prod", "base_url": "https://real-proxy.com"}
|
||||
with patch(
|
||||
"litellm.litellm_core_utils.cli_token_utils.load_cli_token",
|
||||
return_value=token_data,
|
||||
):
|
||||
assert get_stored_api_key(expected_base_url="https://evil.com") is None
|
||||
|
||||
def test_get_stored_api_key_old_token_no_base_url(self):
|
||||
"""Old tokens without a base_url field are rejected when origin check is requested"""
|
||||
token_data = {"key": "sk-old-token"}
|
||||
with patch(
|
||||
"litellm.litellm_core_utils.cli_token_utils.load_cli_token",
|
||||
return_value=token_data,
|
||||
):
|
||||
assert (
|
||||
get_stored_api_key(expected_base_url="https://real-proxy.com") is None
|
||||
)
|
||||
|
||||
|
||||
class TestLoginCommand:
|
||||
"""Test login CLI command"""
|
||||
|
|
|
|||
|
|
@ -220,6 +220,27 @@ class TestToolPermissionGuardrail:
|
|||
assert tool_calls[0].id == "call_123"
|
||||
assert tool_calls[0].function.name == "Read"
|
||||
|
||||
def test_extract_tool_calls_legacy_function_call_format(self):
|
||||
response = ModelResponse(
|
||||
choices=[
|
||||
Choices(
|
||||
message={
|
||||
"function_call": {
|
||||
"name": "Read",
|
||||
"arguments": '{"file_path": "/test/file.txt"}',
|
||||
},
|
||||
}
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
tool_calls = self.guardrail._extract_tool_calls_from_response(response)
|
||||
assert len(tool_calls) == 1
|
||||
assert isinstance(tool_calls[0], ChatCompletionMessageToolCall)
|
||||
assert tool_calls[0].id == "legacy_function_call_0"
|
||||
assert tool_calls[0].function.name == "Read"
|
||||
assert tool_calls[0].function.arguments == '{"file_path": "/test/file.txt"}'
|
||||
|
||||
def test_extract_tool_calls_empty_response(self):
|
||||
response = ModelResponse(choices=[])
|
||||
tool_calls = self.guardrail._extract_tool_calls_from_response(response)
|
||||
|
|
@ -271,6 +292,31 @@ class TestToolPermissionGuardrail:
|
|||
data=data, user_api_key_dict=user_api_key_dict, response=response
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_post_call_success_hook_with_denied_legacy_function_call_raises(
|
||||
self,
|
||||
):
|
||||
response = ModelResponse(
|
||||
choices=[
|
||||
Choices(
|
||||
message={
|
||||
"function_call": {
|
||||
"name": "Read",
|
||||
"arguments": "{}",
|
||||
},
|
||||
}
|
||||
)
|
||||
]
|
||||
)
|
||||
user_api_key_dict = UserAPIKeyAuth()
|
||||
data = {"guardrails": ["test-tool-permission"]}
|
||||
|
||||
with patch.object(self.guardrail, "should_run_guardrail", return_value=True):
|
||||
with pytest.raises(GuardrailRaisedException):
|
||||
await self.guardrail.async_post_call_success_hook(
|
||||
data=data, user_api_key_dict=user_api_key_dict, response=response
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_post_call_success_hook_param_patterns_allow(self):
|
||||
guardrail = ToolPermissionGuardrail(
|
||||
|
|
@ -379,7 +425,9 @@ class TestToolPermissionGuardrail:
|
|||
assert "berri" in choice.message.content
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_post_call_success_hook_missing_arguments_default_allows(self):
|
||||
async def test_async_post_call_success_hook_missing_arguments_blocks_param_rule(
|
||||
self,
|
||||
):
|
||||
guardrail = ToolPermissionGuardrail(
|
||||
guardrail_name="mail-guardrail",
|
||||
rules=[
|
||||
|
|
@ -405,9 +453,52 @@ class TestToolPermissionGuardrail:
|
|||
data = {"guardrails": ["mail-guardrail"]}
|
||||
|
||||
with patch.object(guardrail, "should_run_guardrail", return_value=True):
|
||||
await guardrail.async_post_call_success_hook(
|
||||
data=data, user_api_key_dict=user_api_key_dict, response=response
|
||||
)
|
||||
with pytest.raises(GuardrailRaisedException):
|
||||
await guardrail.async_post_call_success_hook(
|
||||
data=data, user_api_key_dict=user_api_key_dict, response=response
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"arguments",
|
||||
[
|
||||
"{not-json",
|
||||
'["owner@berri.ai"]',
|
||||
],
|
||||
)
|
||||
async def test_async_post_call_success_hook_malformed_arguments_blocks_param_rule(
|
||||
self, arguments
|
||||
):
|
||||
guardrail = ToolPermissionGuardrail(
|
||||
guardrail_name="mail-guardrail",
|
||||
rules=[
|
||||
{
|
||||
"id": "deny_gmail",
|
||||
"tool_name": r"^mail_mcp-send_email$",
|
||||
"decision": "deny",
|
||||
"allowed_param_patterns": {"to[]": r"^.+@gmail\.com$"},
|
||||
}
|
||||
],
|
||||
default_action="allow",
|
||||
on_disallowed_action="block",
|
||||
)
|
||||
|
||||
tool_call = {
|
||||
"function": {
|
||||
"name": "mail_mcp-send_email",
|
||||
"arguments": arguments,
|
||||
},
|
||||
"type": "function",
|
||||
}
|
||||
response = ModelResponse(choices=[Choices(message={"tool_calls": [tool_call]})])
|
||||
user_api_key_dict = UserAPIKeyAuth()
|
||||
data = {"guardrails": ["mail-guardrail"]}
|
||||
|
||||
with patch.object(guardrail, "should_run_guardrail", return_value=True):
|
||||
with pytest.raises(GuardrailRaisedException):
|
||||
await guardrail.async_post_call_success_hook(
|
||||
data=data, user_api_key_dict=user_api_key_dict, response=response
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_pre_call_hook_block_mode(self):
|
||||
|
|
@ -430,6 +521,65 @@ class TestToolPermissionGuardrail:
|
|||
)
|
||||
assert excinfo.value.status_code == 400
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_pre_call_hook_blocks_legacy_functions(self):
|
||||
data = {
|
||||
"functions": [
|
||||
{"name": "Bash", "description": "allowed"},
|
||||
{"name": "Read", "description": "denied"},
|
||||
]
|
||||
}
|
||||
user_api_key_dict = UserAPIKeyAuth()
|
||||
cache = DualCache(default_in_memory_ttl=1)
|
||||
|
||||
with patch.object(self.guardrail, "should_run_guardrail", return_value=True):
|
||||
with pytest.raises(HTTPException) as excinfo:
|
||||
await self.guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=cache,
|
||||
data=data,
|
||||
call_type="completion",
|
||||
)
|
||||
assert excinfo.value.status_code == 400
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_pre_call_hook_blocks_named_legacy_function_call(self):
|
||||
data = {
|
||||
"functions": [{"name": "Bash"}],
|
||||
"function_call": {"name": "Read"},
|
||||
}
|
||||
user_api_key_dict = UserAPIKeyAuth()
|
||||
cache = DualCache(default_in_memory_ttl=1)
|
||||
|
||||
with patch.object(self.guardrail, "should_run_guardrail", return_value=True):
|
||||
with pytest.raises(HTTPException) as excinfo:
|
||||
await self.guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=cache,
|
||||
data=data,
|
||||
call_type="completion",
|
||||
)
|
||||
assert excinfo.value.status_code == 400
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_pre_call_hook_blocks_named_tool_choice(self):
|
||||
data = {
|
||||
"tools": [{"type": "function", "function": {"name": "Bash"}}],
|
||||
"tool_choice": {"type": "function", "function": {"name": "Read"}},
|
||||
}
|
||||
user_api_key_dict = UserAPIKeyAuth()
|
||||
cache = DualCache(default_in_memory_ttl=1)
|
||||
|
||||
with patch.object(self.guardrail, "should_run_guardrail", return_value=True):
|
||||
with pytest.raises(HTTPException) as excinfo:
|
||||
await self.guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=cache,
|
||||
data=data,
|
||||
call_type="completion",
|
||||
)
|
||||
assert excinfo.value.status_code == 400
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_pre_call_hook_uses_custom_template(self):
|
||||
guardrail = ToolPermissionGuardrail(
|
||||
|
|
@ -491,6 +641,41 @@ class TestToolPermissionGuardrail:
|
|||
assert "Bash" in tool_names
|
||||
assert "Read" not in tool_names
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_pre_call_hook_rewrite_mode_filters_legacy_functions(self):
|
||||
guardrail = ToolPermissionGuardrail(
|
||||
guardrail_name="test-tool-permission",
|
||||
rules=self.test_rules,
|
||||
default_action="deny",
|
||||
on_disallowed_action="rewrite",
|
||||
)
|
||||
data = {
|
||||
"functions": [
|
||||
{"name": "Bash", "description": "allowed"},
|
||||
{"name": "Read", "description": "denied"},
|
||||
],
|
||||
"function_call": {"name": "Read"},
|
||||
"tools": [
|
||||
{"type": "function", "function": {"name": "Bash"}},
|
||||
],
|
||||
"tool_choice": {"type": "function", "function": {"name": "Read"}},
|
||||
}
|
||||
user_api_key_dict = UserAPIKeyAuth()
|
||||
cache = DualCache(default_in_memory_ttl=1)
|
||||
|
||||
with patch.object(guardrail, "should_run_guardrail", return_value=True):
|
||||
new_data = await guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=cache,
|
||||
data=data,
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
assert isinstance(new_data, dict)
|
||||
assert [function["name"] for function in new_data["functions"]] == ["Bash"]
|
||||
assert new_data["function_call"] == "none"
|
||||
assert new_data["tool_choice"] == "none"
|
||||
|
||||
def test_modify_response_with_permission_errors(self):
|
||||
# Setup a response with one tool_call
|
||||
tool_call = ChatCompletionMessageToolCall(
|
||||
|
|
@ -522,6 +707,40 @@ class TestToolPermissionGuardrail:
|
|||
assert isinstance(choice.message.content, str)
|
||||
assert "Permission denied" in choice.message.content
|
||||
|
||||
def test_modify_response_with_permission_errors_filters_legacy_function_call(self):
|
||||
response = ModelResponse(
|
||||
choices=[
|
||||
Choices(
|
||||
message={
|
||||
"function_call": {
|
||||
"name": "Read",
|
||||
"arguments": "{}",
|
||||
},
|
||||
"content": "",
|
||||
}
|
||||
)
|
||||
]
|
||||
)
|
||||
tool_call = self.guardrail._extract_tool_calls_from_response(response)[0]
|
||||
denied_tools = [
|
||||
(
|
||||
tool_call,
|
||||
PermissionError(
|
||||
tool_name="Read",
|
||||
rule_id="deny_read",
|
||||
message="Tool 'Read' denied by rule 'deny_read'",
|
||||
),
|
||||
)
|
||||
]
|
||||
|
||||
self.guardrail._modify_response_with_permission_errors(response, denied_tools)
|
||||
|
||||
choice = response.choices[0]
|
||||
assert isinstance(choice, Choices)
|
||||
assert choice.message.function_call is None
|
||||
assert isinstance(choice.message.content, str)
|
||||
assert "Permission denied" in choice.message.content
|
||||
|
||||
|
||||
class TestToolPermissionGuardrailIntegration:
|
||||
"""Integration tests for Tool Permission Guardrail"""
|
||||
|
|
|
|||
|
|
@ -778,3 +778,373 @@ def test_get_callback_identifier_custom_logger_registry_and_fallback():
|
|||
result = get_callback_identifier(my_callback_function)
|
||||
# Should fall back to callback_name() which returns __name__
|
||||
assert result == "my_callback_function"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# /health response shape: model-access scoping and display-field allowlist
|
||||
# ---------------------------------------------------------------------------
|
||||
# These tests pin the contract that the /health response (a) only includes
|
||||
# deployments the calling key is allowed to see, and (b) does not return
|
||||
# provider routing fields like api_base / api_version. They guard against
|
||||
# regressions that would widen the response shape.
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_endpoint_filters_model_list_by_user_access():
|
||||
"""
|
||||
health_endpoint() should restrict _llm_model_list to deployments whose
|
||||
model_name appears in user_api_key_dict.models before running the health
|
||||
check. A key scoped to ["model-a"] should only see model-a in the result,
|
||||
not other deployments configured on the proxy.
|
||||
"""
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.health_endpoints._health_endpoints import health_endpoint
|
||||
|
||||
full_model_list = [
|
||||
{
|
||||
"model_name": "model-a",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o",
|
||||
"api_base": "https://example-a.test",
|
||||
},
|
||||
"model_info": {"id": "id-a"},
|
||||
},
|
||||
{
|
||||
"model_name": "model-b",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o",
|
||||
"api_base": "https://example-b.test",
|
||||
"api_version": "2024-10-21",
|
||||
},
|
||||
"model_info": {"id": "id-b"},
|
||||
},
|
||||
]
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="hashed-test-key",
|
||||
models=["model-a"],
|
||||
)
|
||||
|
||||
captured: dict = {}
|
||||
|
||||
async def fake_perform(**kwargs):
|
||||
captured["model_list"] = kwargs["model_list"]
|
||||
return {
|
||||
"healthy_endpoints": [],
|
||||
"unhealthy_endpoints": [],
|
||||
"healthy_count": 0,
|
||||
"unhealthy_count": 0,
|
||||
}
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.llm_model_list", full_model_list),
|
||||
patch("litellm.proxy.proxy_server.llm_router", None),
|
||||
patch("litellm.proxy.proxy_server.prisma_client", None),
|
||||
patch("litellm.proxy.proxy_server.use_background_health_checks", False),
|
||||
patch("litellm.proxy.proxy_server.user_model", None),
|
||||
patch("litellm.proxy.proxy_server.health_check_results", {}),
|
||||
patch("litellm.proxy.proxy_server.health_check_details", True),
|
||||
patch("litellm.proxy.proxy_server.health_check_concurrency", 1),
|
||||
patch(
|
||||
"litellm.proxy.health_endpoints._health_endpoints._perform_health_check_and_save",
|
||||
side_effect=fake_perform,
|
||||
),
|
||||
):
|
||||
from fastapi import Response
|
||||
|
||||
await health_endpoint(response=Response(), user_api_key_dict=user_api_key_dict)
|
||||
|
||||
assert (
|
||||
"model_list" in captured
|
||||
), "health_endpoint did not call _perform_health_check_and_save"
|
||||
returned_names = {m["model_name"] for m in captured["model_list"]}
|
||||
assert returned_names == {
|
||||
"model-a"
|
||||
}, f"health_endpoint did not scope model_list to caller access: {returned_names}"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_endpoint_filters_background_cache_by_user_access():
|
||||
"""
|
||||
When background_health_checks is enabled, health_endpoint() should also
|
||||
scope the cached result to the caller's allowed models rather than
|
||||
returning the cache verbatim.
|
||||
"""
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.health_endpoints._health_endpoints import health_endpoint
|
||||
|
||||
full_model_list = [
|
||||
{
|
||||
"model_name": "model-a",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o",
|
||||
"api_base": "https://example-a.test",
|
||||
},
|
||||
"model_info": {"id": "id-a"},
|
||||
},
|
||||
{
|
||||
"model_name": "model-b",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o",
|
||||
"api_base": "https://example-b.test",
|
||||
},
|
||||
"model_info": {"id": "id-b"},
|
||||
},
|
||||
]
|
||||
|
||||
cached_results = {
|
||||
"healthy_endpoints": [
|
||||
{
|
||||
"model": "openai/gpt-4o",
|
||||
"model_id": "id-a",
|
||||
"api_base": "https://example-a.test",
|
||||
},
|
||||
{
|
||||
"model": "openai/gpt-4o",
|
||||
"model_id": "id-b",
|
||||
"api_base": "https://example-b.test",
|
||||
},
|
||||
],
|
||||
"unhealthy_endpoints": [],
|
||||
"healthy_count": 2,
|
||||
"unhealthy_count": 0,
|
||||
}
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="hashed-test-key",
|
||||
models=["model-a"],
|
||||
)
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.llm_model_list", full_model_list),
|
||||
patch("litellm.proxy.proxy_server.llm_router", None),
|
||||
patch("litellm.proxy.proxy_server.prisma_client", None),
|
||||
patch("litellm.proxy.proxy_server.use_background_health_checks", True),
|
||||
patch("litellm.proxy.proxy_server.user_model", None),
|
||||
patch("litellm.proxy.proxy_server.health_check_results", cached_results),
|
||||
patch("litellm.proxy.proxy_server.health_check_details", True),
|
||||
patch("litellm.proxy.proxy_server.health_check_concurrency", 1),
|
||||
):
|
||||
from fastapi import Response
|
||||
|
||||
result = await health_endpoint(
|
||||
response=Response(), user_api_key_dict=user_api_key_dict
|
||||
)
|
||||
|
||||
# Sanity: the source cache had two entries before scoping; the scoping
|
||||
# step is what reduces it to one. (This guards against the test passing
|
||||
# vacuously when the cache filter drops everything because cached
|
||||
# entries lack the model_id key — both entries carry model_id above.)
|
||||
assert len(cached_results["healthy_endpoints"]) == 2
|
||||
assert all(
|
||||
ep.get("model_id") for ep in cached_results["healthy_endpoints"]
|
||||
), "test fixture invariant: every cached entry must carry a model_id"
|
||||
|
||||
# The non-admin caller must not see api_base on the returned cache entries.
|
||||
returned = result.get("healthy_endpoints", [])
|
||||
assert (
|
||||
len(returned) == 1
|
||||
), f"expected exactly one cached entry after scoping, got {len(returned)}"
|
||||
assert returned[0]["model_id"] == "id-a"
|
||||
assert "api_base" not in returned[0]
|
||||
assert result["healthy_count"] == 1
|
||||
assert result["unhealthy_count"] == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_endpoint_admin_sees_routing_fields_non_admin_does_not():
|
||||
"""
|
||||
A proxy admin should still see ``api_base`` and ``api_version`` in the
|
||||
/health response so they can tell which Vertex region / Azure resource
|
||||
+ API version is healthy. A non-admin caller must not — both fields
|
||||
should be stripped, and the response should carry a notice header so
|
||||
non-admin clients can detect the change programmatically.
|
||||
"""
|
||||
from fastapi import Response
|
||||
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.health_endpoints._health_endpoints import health_endpoint
|
||||
|
||||
full_model_list = [
|
||||
{
|
||||
"model_name": "model-a",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o",
|
||||
"api_base": "https://example-a.test",
|
||||
},
|
||||
"model_info": {"id": "id-a"},
|
||||
},
|
||||
]
|
||||
cached_results = {
|
||||
"healthy_endpoints": [
|
||||
{
|
||||
"model": "openai/gpt-4o",
|
||||
"model_id": "id-a",
|
||||
"api_base": "https://us-central1-aiplatform.googleapis.com/v1/projects/p",
|
||||
"api_version": "2024-10-21",
|
||||
},
|
||||
],
|
||||
"unhealthy_endpoints": [],
|
||||
"healthy_count": 1,
|
||||
"unhealthy_count": 0,
|
||||
}
|
||||
|
||||
admin_key = UserAPIKeyAuth(
|
||||
api_key="hashed-admin-key",
|
||||
models=["model-a"],
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
)
|
||||
non_admin_key = UserAPIKeyAuth(
|
||||
api_key="hashed-user-key",
|
||||
models=["model-a"],
|
||||
)
|
||||
|
||||
common_patches = [
|
||||
patch("litellm.proxy.proxy_server.llm_model_list", full_model_list),
|
||||
patch("litellm.proxy.proxy_server.llm_router", None),
|
||||
patch("litellm.proxy.proxy_server.prisma_client", None),
|
||||
patch("litellm.proxy.proxy_server.use_background_health_checks", True),
|
||||
patch("litellm.proxy.proxy_server.user_model", None),
|
||||
patch("litellm.proxy.proxy_server.health_check_results", cached_results),
|
||||
patch("litellm.proxy.proxy_server.health_check_details", True),
|
||||
patch("litellm.proxy.proxy_server.health_check_concurrency", 1),
|
||||
]
|
||||
|
||||
for p in common_patches:
|
||||
p.start()
|
||||
try:
|
||||
admin_response = Response()
|
||||
non_admin_response = Response()
|
||||
admin_result = await health_endpoint(
|
||||
response=admin_response, user_api_key_dict=admin_key
|
||||
)
|
||||
non_admin_result = await health_endpoint(
|
||||
response=non_admin_response, user_api_key_dict=non_admin_key
|
||||
)
|
||||
finally:
|
||||
for p in common_patches:
|
||||
p.stop()
|
||||
|
||||
admin_eps = admin_result.get("healthy_endpoints", [])
|
||||
non_admin_eps = non_admin_result.get("healthy_endpoints", [])
|
||||
|
||||
assert len(admin_eps) == 1
|
||||
assert (
|
||||
admin_eps[0]["api_base"]
|
||||
== "https://us-central1-aiplatform.googleapis.com/v1/projects/p"
|
||||
), "admin must see the full api_base so they can identify the region"
|
||||
assert (
|
||||
admin_eps[0]["api_version"] == "2024-10-21"
|
||||
), "admin must see api_version so they can distinguish provider deployments"
|
||||
|
||||
assert len(non_admin_eps) == 1
|
||||
assert "api_base" not in non_admin_eps[0]
|
||||
assert "api_version" not in non_admin_eps[0]
|
||||
|
||||
# Non-admin response must advertise that api_base/api_version were
|
||||
# withheld so clients that previously parsed them can detect the change.
|
||||
assert (
|
||||
non_admin_response.headers.get("Litellm-Health-Field-Notice")
|
||||
== "api_base and api_version are admin-only on this endpoint"
|
||||
)
|
||||
assert "Litellm-Health-Field-Notice" not in admin_response.headers
|
||||
|
||||
# Stripping must produce a copy — the shared cache must still carry the
|
||||
# routing fields so the next admin caller can read them.
|
||||
cached_first = cached_results["healthy_endpoints"][0]
|
||||
assert (
|
||||
cached_first["api_base"]
|
||||
== "https://us-central1-aiplatform.googleapis.com/v1/projects/p"
|
||||
)
|
||||
assert cached_first["api_version"] == "2024-10-21"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_endpoint_warns_when_scoped_models_lack_model_id():
|
||||
"""
|
||||
When a scoped key's accessible models exist on the proxy but none of the
|
||||
matching deployments expose a ``model_info.id``, the cache filter drops
|
||||
everything. The response should include a structured ``warnings`` field
|
||||
so the caller can distinguish "no deployments configured" from
|
||||
"deployments excluded due to missing model_info.id".
|
||||
"""
|
||||
from fastapi import Response
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.health_endpoints._health_endpoints import health_endpoint
|
||||
|
||||
full_model_list = [
|
||||
{
|
||||
"model_name": "model-a",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o",
|
||||
"api_base": "https://example-a.test",
|
||||
},
|
||||
# Intentionally no model_info.id — this is the misconfiguration
|
||||
# the warnings field is meant to flag.
|
||||
"model_info": {},
|
||||
},
|
||||
]
|
||||
cached_results = {
|
||||
"healthy_endpoints": [
|
||||
{
|
||||
"model": "openai/gpt-4o",
|
||||
"model_id": "id-a",
|
||||
"api_base": "https://example-a.test",
|
||||
},
|
||||
],
|
||||
"unhealthy_endpoints": [],
|
||||
"healthy_count": 1,
|
||||
"unhealthy_count": 0,
|
||||
}
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="hashed-user-key",
|
||||
models=["model-a"],
|
||||
)
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.llm_model_list", full_model_list),
|
||||
patch("litellm.proxy.proxy_server.llm_router", None),
|
||||
patch("litellm.proxy.proxy_server.prisma_client", None),
|
||||
patch("litellm.proxy.proxy_server.use_background_health_checks", True),
|
||||
patch("litellm.proxy.proxy_server.user_model", None),
|
||||
patch("litellm.proxy.proxy_server.health_check_results", cached_results),
|
||||
patch("litellm.proxy.proxy_server.health_check_details", True),
|
||||
patch("litellm.proxy.proxy_server.health_check_concurrency", 1),
|
||||
):
|
||||
result = await health_endpoint(
|
||||
response=Response(), user_api_key_dict=user_api_key_dict
|
||||
)
|
||||
|
||||
assert result["healthy_count"] == 0
|
||||
assert result["unhealthy_count"] == 0
|
||||
assert "warnings" in result, (
|
||||
"empty cache result must surface a warnings field so the caller "
|
||||
"can distinguish 'no deployments' from 'deployments excluded'"
|
||||
)
|
||||
assert any("model_info.id" in w for w in result["warnings"])
|
||||
|
||||
|
||||
def test_clean_endpoint_data_strips_credentials_keeps_routing_fields():
|
||||
"""
|
||||
_clean_endpoint_data() drops credentials but leaves api_base /
|
||||
api_version intact — the per-caller hide/show happens in the endpoint
|
||||
layer based on user role, not in the cleaning helper. This guarantees
|
||||
proxy admins continue to see those fields in the /health response.
|
||||
"""
|
||||
from litellm.proxy.health_check import _clean_endpoint_data
|
||||
|
||||
raw = {
|
||||
"model": "openai/gpt-4o",
|
||||
"api_key": "sk-test",
|
||||
"api_base": "https://example.test/v1",
|
||||
"api_version": "2024-10-21",
|
||||
"aws_access_key_id": "AKIAEXAMPLE",
|
||||
}
|
||||
|
||||
cleaned = _clean_endpoint_data(raw, details=True)
|
||||
|
||||
assert "api_key" not in cleaned
|
||||
assert "aws_access_key_id" not in cleaned
|
||||
assert cleaned.get("api_base") == "https://example.test/v1"
|
||||
assert cleaned.get("api_version") == "2024-10-21"
|
||||
|
|
|
|||
118
tests/test_litellm/proxy/test_lazy_openapi_snapshot.py
Normal file
118
tests/test_litellm/proxy/test_lazy_openapi_snapshot.py
Normal file
|
|
@ -0,0 +1,118 @@
|
|||
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}": {
|
||||
"delete": {"operationId": "proxy_route_proxy__endpoint__put"},
|
||||
"get": {"operationId": "proxy_route_proxy__endpoint__put"},
|
||||
"post": {"operationId": "proxy_route_proxy__endpoint__put"},
|
||||
"put": {"operationId": "proxy_route_proxy__endpoint__put"},
|
||||
}
|
||||
}
|
||||
|
||||
_normalize_operation_ids(paths)
|
||||
|
||||
operations = paths["/proxy/{endpoint}"]
|
||||
assert operations["delete"]["operationId"] == "proxy_route_proxy__endpoint__delete"
|
||||
assert operations["get"]["operationId"] == "proxy_route_proxy__endpoint__get"
|
||||
assert operations["post"]["operationId"] == "proxy_route_proxy__endpoint__post"
|
||||
assert operations["put"]["operationId"] == "proxy_route_proxy__endpoint__put"
|
||||
|
||||
|
||||
def test_normalize_operation_ids_preserves_custom_ids():
|
||||
paths = {
|
||||
"/proxy/{endpoint}": {
|
||||
"get": {"operationId": "custom_operation"},
|
||||
"post": {"operationId": "custom_operation"},
|
||||
}
|
||||
}
|
||||
|
||||
_normalize_operation_ids(paths)
|
||||
|
||||
operations = paths["/proxy/{endpoint}"]
|
||||
assert operations["get"]["operationId"] == "custom_operation"
|
||||
assert operations["post"]["operationId"] == "custom_operation"
|
||||
|
|
@ -5988,6 +5988,202 @@ async def test_reseed_warms_cache_even_on_zero_db_spend():
|
|||
ps.prisma_client = orig_prisma
|
||||
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# /config/update — critical paths only.
|
||||
#
|
||||
# These exercise the four behaviors that broke or changed in the rewrite of
|
||||
# update_config (litellm/proxy/proxy_server.py): targeted per-section writes,
|
||||
# the removal of the store_model_in_db gate, env var encryption, and the
|
||||
# success_callback / litellm_settings merge semantics. All other branches
|
||||
# (auth, missing-DB, slack auto-enable, router_settings merge) are covered
|
||||
# implicitly or by upstream tests.
|
||||
# -----------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _FakeRow:
|
||||
def __init__(self, param_name, param_value):
|
||||
self.param_name = param_name
|
||||
self.param_value = param_value
|
||||
|
||||
|
||||
class _FakeLitellmConfig:
|
||||
def __init__(self, initial_rows=None):
|
||||
self.rows = dict(initial_rows or {})
|
||||
self.upsert_calls: list = []
|
||||
self.find_first = AsyncMock(side_effect=self._find_first)
|
||||
self.upsert = AsyncMock(side_effect=self._upsert)
|
||||
|
||||
async def _find_first(self, where=None):
|
||||
if where and "param_name" in where:
|
||||
name = where["param_name"]
|
||||
if name in self.rows:
|
||||
return _FakeRow(name, self.rows[name])
|
||||
return None
|
||||
|
||||
async def _upsert(self, where=None, data=None):
|
||||
name = where["param_name"]
|
||||
raw = data["update"]["param_value"]
|
||||
value = json.loads(raw) if isinstance(raw, str) else raw
|
||||
self.rows[name] = value
|
||||
self.upsert_calls.append((name, value))
|
||||
|
||||
|
||||
class _FakePrismaClient:
|
||||
def __init__(self, initial_rows=None):
|
||||
self.db = mock.MagicMock()
|
||||
self.db.litellm_config = _FakeLitellmConfig(initial_rows=initial_rows)
|
||||
self.jsonify_object = lambda obj: obj
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def _update_config_setup(monkeypatch):
|
||||
"""Install fakes for the /config/update endpoint and return (client, prisma)."""
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth as auth_dep
|
||||
|
||||
def _install(initial_rows=None, store_model_in_db=True):
|
||||
prisma = _FakePrismaClient(initial_rows=initial_rows)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.store_model_in_db", store_model_in_db
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.encrypt_value_helper",
|
||||
lambda value, **_: f"enc:{value}",
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.invalidate_config_param",
|
||||
AsyncMock(return_value=None),
|
||||
)
|
||||
from litellm.proxy.proxy_server import proxy_config as real_proxy_config
|
||||
|
||||
monkeypatch.setattr(
|
||||
real_proxy_config, "add_deployment", AsyncMock(return_value=None)
|
||||
)
|
||||
|
||||
original_overrides = app.dependency_overrides.copy()
|
||||
app.dependency_overrides[auth_dep] = lambda: UserAPIKeyAuth(
|
||||
user_id="test_admin",
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-1234",
|
||||
)
|
||||
client = TestClient(app)
|
||||
|
||||
def _restore():
|
||||
app.dependency_overrides = original_overrides
|
||||
|
||||
return client, prisma, _restore
|
||||
|
||||
return _install
|
||||
|
||||
|
||||
def test_update_config_writes_only_sent_section(_update_config_setup):
|
||||
"""A request that only touches general_settings must not write any other
|
||||
section row, and must leave previously-written rows byte-identical."""
|
||||
client, prisma, restore = _update_config_setup(
|
||||
initial_rows={
|
||||
"litellm_settings": {"drop_params": True},
|
||||
"environment_variables": {"FOO": "enc:bar"},
|
||||
}
|
||||
)
|
||||
try:
|
||||
resp = client.post(
|
||||
"/config/update",
|
||||
json={"general_settings": {"store_prompts_in_spend_logs": True}},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
written = {name for name, _ in prisma.db.litellm_config.upsert_calls}
|
||||
assert written == {"general_settings"}
|
||||
assert prisma.db.litellm_config.rows["litellm_settings"] == {
|
||||
"drop_params": True
|
||||
}
|
||||
assert prisma.db.litellm_config.rows["environment_variables"] == {
|
||||
"FOO": "enc:bar"
|
||||
}
|
||||
finally:
|
||||
restore()
|
||||
|
||||
|
||||
def test_update_config_can_flip_store_model_in_db_when_currently_false(
|
||||
_update_config_setup,
|
||||
):
|
||||
"""The endpoint used to refuse all writes when store_model_in_db was
|
||||
False, blocking the very request that would flip it to True."""
|
||||
client, prisma, restore = _update_config_setup(store_model_in_db=False)
|
||||
try:
|
||||
resp = client.post(
|
||||
"/config/update", json={"general_settings": {"store_model_in_db": True}}
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
assert (
|
||||
prisma.db.litellm_config.rows["general_settings"]["store_model_in_db"]
|
||||
is True
|
||||
)
|
||||
finally:
|
||||
restore()
|
||||
|
||||
|
||||
def test_update_config_environment_variables_encrypted_before_write(
|
||||
_update_config_setup,
|
||||
):
|
||||
"""env var values must be encrypted before they hit the DB row."""
|
||||
client, prisma, restore = _update_config_setup()
|
||||
try:
|
||||
resp = client.post(
|
||||
"/config/update",
|
||||
json={"environment_variables": {"OPENAI_API_KEY": "sk-secret"}},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
stored = prisma.db.litellm_config.rows["environment_variables"]
|
||||
assert stored == {"OPENAI_API_KEY": "enc:sk-secret"}
|
||||
finally:
|
||||
restore()
|
||||
|
||||
|
||||
def test_update_config_litellm_settings_request_wins_for_non_callback_keys(
|
||||
_update_config_setup,
|
||||
):
|
||||
"""Sending {"drop_params": False} when the row holds drop_params: True
|
||||
must persist False (request wins). Untouched keys preserved."""
|
||||
client, prisma, restore = _update_config_setup(
|
||||
initial_rows={
|
||||
"litellm_settings": {"drop_params": True, "set_verbose": True},
|
||||
}
|
||||
)
|
||||
try:
|
||||
resp = client.post(
|
||||
"/config/update", json={"litellm_settings": {"drop_params": False}}
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
stored = prisma.db.litellm_config.rows["litellm_settings"]
|
||||
assert stored["drop_params"] is False
|
||||
assert stored["set_verbose"] is True
|
||||
finally:
|
||||
restore()
|
||||
|
||||
|
||||
def test_update_config_success_callback_normalizes_existing_mixed_case(
|
||||
_update_config_setup,
|
||||
):
|
||||
"""Existing mixed-case callback names (written elsewhere) must be
|
||||
normalized to lowercase before union, otherwise the union dedup misses
|
||||
against the lowercase incoming entry and delete_callback (lowercase
|
||||
lookup) cannot find the original."""
|
||||
client, prisma, restore = _update_config_setup(
|
||||
initial_rows={"litellm_settings": {"success_callback": ["Langfuse", "SQS"]}}
|
||||
)
|
||||
try:
|
||||
resp = client.post(
|
||||
"/config/update",
|
||||
json={"litellm_settings": {"success_callback": ["langfuse"]}},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
stored = prisma.db.litellm_config.rows["litellm_settings"]["success_callback"]
|
||||
assert set(stored) == {"langfuse", "sqs"}
|
||||
finally:
|
||||
restore()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Lazy feature loading (LazyFeatureMiddleware) — verifies that optional
|
||||
# routers are NOT imported at module load and ARE imported on first request
|
||||
|
|
@ -5996,9 +6192,6 @@ async def test_reseed_warms_cache_even_on_zero_db_spend():
|
|||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
import sys
|
||||
|
||||
|
||||
class TestLazyFeatureRegistry:
|
||||
"""Sanity checks on the registry shape — guards against accidental edits."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -6,8 +6,8 @@ import {
|
|||
deriveErrorMessage,
|
||||
handleError,
|
||||
} from "@/components/networking";
|
||||
import { all_admin_roles } from "@/utils/roles";
|
||||
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
|
||||
import { all_admin_roles } from "@/utils/roles";
|
||||
|
||||
// ── Types ────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
|
@ -81,7 +81,6 @@ export const useProjects = () => {
|
|||
return useQuery<ProjectResponse[]>({
|
||||
queryKey: projectKeys.list({}),
|
||||
queryFn: async () => fetchProjects(accessToken!),
|
||||
enabled:
|
||||
Boolean(accessToken) && all_admin_roles.includes(userRole || ""),
|
||||
enabled: Boolean(accessToken) && all_admin_roles.includes(userRole!),
|
||||
});
|
||||
};
|
||||
|
|
|
|||
|
|
@ -169,8 +169,8 @@ const UsagePage: React.FC<UsagePageProps> = ({ teams, organizations }) => {
|
|||
}
|
||||
}, [isAdmin, userID]);
|
||||
|
||||
// For non-admins, always pass their own user_id
|
||||
const effectiveUserId = isAdmin ? selectedUserId : userID || null;
|
||||
// For non-admins or "my-usage" view, always pass their own user_id
|
||||
const effectiveUserId = usageView === "my-usage" || !isAdmin ? userID || null : selectedUserId;
|
||||
|
||||
const startTime = useMemo(() => (dateValue.from ? new Date(dateValue.from) : null), [dateValue.from]);
|
||||
const endTime = useMemo(() => (dateValue.to ? new Date(dateValue.to) : null), [dateValue.to]);
|
||||
|
|
@ -477,10 +477,10 @@ const UsagePage: React.FC<UsagePageProps> = ({ teams, organizations }) => {
|
|||
}
|
||||
/>
|
||||
)}
|
||||
{/* Your Usage Panel */}
|
||||
{usageView === "global" && (
|
||||
{/* Your Usage / Global Usage Panel */}
|
||||
{(usageView === "global" || usageView === "my-usage") && (
|
||||
<>
|
||||
{isAdmin && (
|
||||
{isAdmin && usageView === "global" && (
|
||||
<div className="mb-4">
|
||||
<Text className="mb-2">Filter by user</Text>
|
||||
<Select
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ import {
|
|||
} from "@ant-design/icons";
|
||||
import { Badge, Select } from "antd";
|
||||
import React from "react";
|
||||
export type UsageOption = "global" | "organization" | "team" | "customer" | "tag" | "agent" | "user" | "user-agent-activity";
|
||||
export type UsageOption = "global" | "my-usage" | "organization" | "team" | "customer" | "tag" | "agent" | "user" | "user-agent-activity";
|
||||
export interface UsageViewSelectProps {
|
||||
value: UsageOption;
|
||||
onChange: (value: UsageOption) => void;
|
||||
|
|
@ -43,6 +43,13 @@ const OPTIONS: OptionConfig[] = [
|
|||
descriptionForNonAdmin: "View your usage",
|
||||
icon: <GlobalOutlined style={{ fontSize: "16px" }} />,
|
||||
},
|
||||
{
|
||||
value: "my-usage",
|
||||
label: "Your Usage",
|
||||
description: "View your own usage",
|
||||
icon: <UserOutlined style={{ fontSize: "16px" }} />,
|
||||
adminOnly: true,
|
||||
},
|
||||
{
|
||||
value: "organization",
|
||||
label: "Organization Usage",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue