Litellm ishaan march 20 (#24303)

* feat(redis): add circuit breaker to RedisCache to fast-fail when Redis is down (#24181)

* feat(redis): add circuit breaker env var constants

* feat(redis): add RedisCircuitBreaker and apply guard decorator to all async ops

* fix(dual_cache): fall back to L1 instead of re-raising on Redis increment failures

* test(caching): add circuit breaker unit tests

* fix(redis): fast-fail concurrent HALF_OPEN probes — only one probe at a time

* fix(dual_cache): return None fallback when in_memory_cache is absent and Redis fails

* test(caching): add regression tests for HALF_OPEN concurrency and None fallback

* Fix blocking sync next in __anext__ (#24177)

* Fix blocking sync next

* Update tests/test_litellm/litellm_core_utils/test_streaming_handler.py

Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com>

* fix PEP 479 regression in __anext__ sync iterator exhaustion

asyncio.to_thread re-raises thread exceptions inside a coroutine, where
PEP 479 converts StopIteration to RuntimeError before any except clause
can catch it. Add _next_sync_or_exhausted() module-level helper that
catches StopIteration in the thread and returns a sentinel instead, then
raise StopAsyncIteration in the coroutine.

Also rewrites the non-blocking test to use asyncio.gather() instead of
asyncio.create_task() (which returned None on Python 3.9 / pytest-asyncio
in CI), and adds an exhaustion regression test that drains the wrapper
fully and asserts no RuntimeError leaks out.

---------

Co-authored-by: Emerson Gomes <emerson.gomes@thalesgroup.com>
Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com>

* feat: add git-subdir source type to claude-code/plugins API (#24223)

Support a third plugin source type `git-subdir` alongside the existing
`github` and `url` types, as documented in the official Claude Code
plugin marketplaces spec.

New format: {"source": "git-subdir", "url": "...", "path": "subdir/path"}

- Validates url and path fields are present and non-empty
- Rejects absolute paths, '..' segments, backslashes, and percent-encoded
  traversal sequences (including double-encoded variants via regex check)
- Extracts path validation into _validate_git_subdir_path() helper
- Updates Pydantic field description to document all three source types
- Adds isValidUrl() check for url/git-subdir source types in the UI form
- Adds "Git Subdir" option to the UI form with a required Path field
- Adds unit tests covering success, update, missing/empty fields,
  path traversal variants, and unknown source type

Co-authored-by: Claude Sonnet 4.6 <noreply@anthropic.com>

* [FEAT] add extract_header and extract_footer to Mistral OCR supported params (#24213)

* docs: add git-subdir source type to claude-code plugin marketplace docs (#24289)

* fix(ui): swap J/K keyboard navigation in log details drawer (#24279) (#24286)

J should navigate down (next) and K should navigate up (previous),
matching vim/standard conventions.

* fix: use async_set_cache in user_api_key_auth hot path (#24302)

* fix: use async_set_cache in auth hot path to avoid blocking event loop

* test: assert no blocking set_cache call in _user_api_key_auth_builder

* test: broaden blocking call check to all sync DualCache methods

* test: fix regression test to actually catch blocking cache calls

* fix: ruff lint unused variable + UI build MessageManager error

- litellm/caching/redis_cache.py: remove unused variable 'e' in circuit
  breaker exception handler (F841)
- add_plugin_form.tsx: use MessageManager.error() instead of undefined
  message.error() for git URL validation

Co-authored-by: Ishaan Jaff <ishaan-jaff@users.noreply.github.com>

* docs: add REDIS_CIRCUIT_BREAKER env vars to config_settings reference

Add REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD and
REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT to the environment variables
reference table so test_env_keys.py passes.

Co-authored-by: Ishaan Jaff <ishaan-jaff@users.noreply.github.com>

---------

Co-authored-by: Emerson Gomes <emerson.gomes@thalesgroup.com>
Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com>
Co-authored-by: Vincenzo Barrea <manamana88@users.noreply.github.com>
Co-authored-by: Claude Sonnet 4.6 <noreply@anthropic.com>
Co-authored-by: Robert Kirscht <rkirscht242@gmail.com>
Co-authored-by: Imgyu Kim <kimimgo@gmail.com>
Co-authored-by: Cursor Agent <cursoragent@cursor.com>
Co-authored-by: Ishaan Jaff <ishaan-jaff@users.noreply.github.com>
This commit is contained in:
Ishaan Jaff 2026-03-21 12:40:11 -07:00 • committed by GitHub
parent b64b0d4b9b
commit 2ea9e207bd
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
21 changed files with 1245 additions and 110 deletions

View file

@ -952,6 +952,8 @@ router_settings:
| QDRANT_URL | Connection URL for Qdrant database
| QDRANT_VECTOR_SIZE | Vector size for Qdrant operations. Default is 1536
| REDIS_CONNECTION_POOL_TIMEOUT | Timeout in seconds for Redis connection pool. Default is 5
| REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD | Number of consecutive failures before the Redis circuit breaker opens. Default is 5
| REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT | Time in seconds before the Redis circuit breaker attempts recovery after opening. Default is 60
| REDIS_CLUSTER_NODES | JSON-formatted list of Redis cluster startup nodes for Redis Cluster mode. Example: `[{"host": "node1", "port": 6379}]`
| REDIS_HOST | Hostname for Redis server
| REDIS_PASSWORD | Password for Redis service

View file

@ -37,7 +37,7 @@ Click **+ Add New Plugin** to register a plugin in your marketplace.
Enter the plugin information:
- **Name**: Plugin identifier (kebab-case, e.g., `my-plugin`)
- **Source Type**: Choose GitHub or URL
- **Source Type**: Choose GitHub, Git URL, or Git Subdir
- **Repository/URL**: The git source (e.g., `org/repo` for GitHub)
- **Version**: Semantic version (optional)
- **Description**: What the plugin does
@ -216,6 +216,22 @@ curl -X DELETE http://localhost:4000/claude-code/plugins/my-plugin \
Use this format for GitLab, Bitbucket, or self-hosted git repositories.
</TabItem>
<TabItem value="git-subdir" label="Git Subdir">
```json
{
"name": "my-plugin",
"source": {
"source": "git-subdir",
"url": "https://github.com/org/repo.git",
"path": "plugins/my-plugin"
}
}
```
Use this format when your plugin lives in a subdirectory of a git repository. The `path` field must be a relative path of slash-separated segments (alphanumeric, dots, hyphens, underscores only).
</TabItem>
</Tabs>

View file

@ -393,16 +393,17 @@ class DualCache(BaseCache):
parent_otel_span: Optional[Span] = None,
local_only: bool = False,
**kwargs,
) -> float:
) -> Optional[float]:
"""
Key - the key in cache
Value - float - the value you want to increment by
Returns - float - the incremented value
Returns - the incremented value, or None if no cache backend is
available (in_memory_cache is None and Redis failed/is absent).
"""
result: Optional[float] = None
try:
result: float = value
if self.in_memory_cache is not None:
result = await self.in_memory_cache.async_increment(
key, value, **kwargs
@ -418,7 +419,11 @@ class DualCache(BaseCache):
return result
except Exception as e:
raise e # don't log if exception is raised
verbose_logger.warning(
"Redis async_increment_cache failed, falling back to in-memory result: %s",
e,
)
return result
async def async_increment_cache_pipeline(
self,
@ -427,8 +432,8 @@ class DualCache(BaseCache):
parent_otel_span: Optional[Span] = None,
**kwargs,
) -> Optional[List[float]]:
result: Optional[List[float]] = None
try:
result: Optional[List[float]] = None
if self.in_memory_cache is not None:
result = await self.in_memory_cache.async_increment_pipeline(
increment_list=increment_list,
@ -443,7 +448,11 @@ class DualCache(BaseCache):
return result
except Exception as e:
raise e # don't log if exception is raised
verbose_logger.warning(
"Redis async_increment_cache_pipeline failed, falling back to in-memory result: %s",
e,
)
return result
async def async_set_cache_sadd(
self, key, value: List, local_only: bool = False, **kwargs

View file

@ -10,6 +10,7 @@ Has 4 primary methods:
import ast
import asyncio
import functools
import hashlib
import inspect
import json
@ -19,7 +20,11 @@ from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union, cast
import litellm
from litellm._logging import print_verbose, verbose_logger
from litellm.constants import DEFAULT_REDIS_MAJOR_VERSION
from litellm.constants import (
DEFAULT_REDIS_MAJOR_VERSION,
REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD,
REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT,
)
from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs
from litellm.litellm_core_utils.coroutine_checker import coroutine_checker
from litellm.types.caching import (
@ -89,6 +94,91 @@ def _get_call_stack_info(num_frames: int = 2) -> str:
return "unknown"
class RedisCircuitBreaker:
"""
Tracks Redis health for a RedisCache instance.
States:
CLOSED - normal, Redis is called
OPEN - Redis is down, raise immediately (no network call)
HALF_OPEN - recovery probe: allow one request through
Transitions:
CLOSED -> OPEN after failure_threshold consecutive failures
OPEN -> HALF_OPEN after recovery_timeout seconds
HALF_OPEN -> CLOSED on success
HALF_OPEN -> OPEN on failure (resets timer)
"""
CLOSED = "closed"
OPEN = "open"
HALF_OPEN = "half_open"
def __init__(self, failure_threshold: int, recovery_timeout: int) -> None:
self.failure_threshold = failure_threshold
self.recovery_timeout = recovery_timeout
self._failure_count = 0
self._opened_at: Optional[float] = None
self._state = self.CLOSED
def is_open(self) -> bool:
"""Returns True if Redis calls should be skipped."""
if self._state == self.HALF_OPEN:
# Probe already in flight — fast-fail all concurrent requests.
# Only the one call that caused the OPEN→HALF_OPEN transition
# (which returned False) is the designated probe.
return True
if self._state == self.OPEN:
if time.time() - (self._opened_at or 0) > self.recovery_timeout:
self._state = self.HALF_OPEN
return False # this caller is the designated probe
return True
return False
def record_failure(self) -> None:
self._failure_count += 1
self._opened_at = time.time()
if self._failure_count >= self.failure_threshold:
if self._state != self.OPEN:
verbose_logger.warning(
"Redis circuit breaker OPENED after %d consecutive failures — "
"fast-failing Redis calls for %ds",
self._failure_count,
self.recovery_timeout,
)
self._state = self.OPEN
def record_success(self) -> None:
if self._state == self.HALF_OPEN:
verbose_logger.info("Redis circuit breaker CLOSED — Redis recovered")
self._failure_count = 0
self._state = self.CLOSED
def _redis_circuit_breaker_guard(method): # type: ignore
"""
Decorator for RedisCache async methods.
Checks the circuit breaker before each call; records success/failure after.
Does not apply to ping/disconnect/test_connection (health/teardown must always run).
"""
@functools.wraps(method)
async def wrapper(self, *args, **kwargs): # type: ignore
if self._circuit_breaker.is_open():
raise Exception(
f"Redis circuit breaker is open — skipping {method.__name__}"
)
try:
result = await method(self, *args, **kwargs)
self._circuit_breaker.record_success()
return result
except Exception:
self._circuit_breaker.record_failure()
raise
return wrapper
class RedisCache(BaseCache):
# if users don't provider one, use the default litellm cache
@ -150,6 +240,11 @@ class RedisCache(BaseCache):
except Exception:
pass
self._circuit_breaker = RedisCircuitBreaker(
failure_threshold=REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD,
recovery_timeout=REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT,
)
self._setup_health_pings()
if litellm.default_redis_ttl is not None:
@ -375,6 +470,7 @@ class RedisCache(BaseCache):
)
raise e
@_redis_circuit_breaker_guard
async def async_scan_iter(self, pattern: str, count: int = 100) -> list:
start_time = time.time()
try:
@ -451,6 +547,7 @@ class RedisCache(BaseCache):
verbose_logger.error(f"Error registering Redis script: {str(e)}")
raise e
@_redis_circuit_breaker_guard
async def async_set_cache(self, key, value, **kwargs):
from redis.asyncio import Redis
@ -560,6 +657,7 @@ class RedisCache(BaseCache):
results = await pipe.execute()
return results
@_redis_circuit_breaker_guard
async def async_set_cache_pipeline(
self, cache_list: List[Tuple[Any, Any]], ttl: Optional[float] = None, **kwargs
):
@ -636,6 +734,7 @@ class RedisCache(BaseCache):
except Exception:
raise
@_redis_circuit_breaker_guard
async def async_set_cache_sadd(
self, key, value: List, ttl: Optional[float], **kwargs
):
@ -708,6 +807,7 @@ class RedisCache(BaseCache):
value,
)
@_redis_circuit_breaker_guard
async def batch_cache_write(self, key, value, **kwargs):
print_verbose(
f"in batch cache writing for redis buffer size={len(self.redis_batch_writing_buffer)}",
@ -717,6 +817,7 @@ class RedisCache(BaseCache):
if len(self.redis_batch_writing_buffer) >= self.redis_flush_size:
await self.flush_cache_buffer() # logging done in here
@_redis_circuit_breaker_guard
async def async_increment(
self,
key,
@ -894,6 +995,7 @@ class RedisCache(BaseCache):
verbose_logger.error(f"Error occurred in batch get cache - {str(e)}")
return key_value_dict
@_redis_circuit_breaker_guard
async def async_get_cache(
self, key, parent_otel_span: Optional[Span] = None, **kwargs
):
@ -944,6 +1046,7 @@ class RedisCache(BaseCache):
f"litellm.caching.caching: async get() - Got exception from REDIS: {str(e)}"
)
@_redis_circuit_breaker_guard
async def async_batch_get_cache(
self,
key_list: Union[List[str], List[Optional[str]]],
@ -1087,6 +1190,7 @@ class RedisCache(BaseCache):
)
raise e
@_redis_circuit_breaker_guard
async def delete_cache_keys(self, keys):
# typed as Any, redis python lib has incomplete type stubs for RedisCluster and does not include `delete`
_redis_client: Any = self.init_async_client()
@ -1151,6 +1255,7 @@ class RedisCache(BaseCache):
"error": str(e),
}
@_redis_circuit_breaker_guard
async def async_delete_cache(self, key: str):
# typed as Any, redis python lib has incomplete type stubs for RedisCluster and does not include `delete`
_redis_client: Any = self.init_async_client()
@ -1184,6 +1289,7 @@ class RedisCache(BaseCache):
)
return [r for r in results if isinstance(r, float)]
@_redis_circuit_breaker_guard
async def async_increment_pipeline(
self, increment_list: List[RedisPipelineIncrementOperation], **kwargs
) -> Optional[List[float]]:
@ -1247,6 +1353,7 @@ class RedisCache(BaseCache):
)
raise e
@_redis_circuit_breaker_guard
async def async_get_ttl(self, key: str) -> Optional[int]:
"""
Get the remaining TTL of a key in Redis
@ -1270,6 +1377,7 @@ class RedisCache(BaseCache):
verbose_logger.debug(f"Redis TTL Error: {e}")
return None
@_redis_circuit_breaker_guard
async def async_rpush(
self,
key: str,
@ -1336,6 +1444,7 @@ class RedisCache(BaseCache):
raise r
return results
@_redis_circuit_breaker_guard
async def async_rpush_pipeline(
self,
rpush_list: List[RedisPipelineRpushOperation],
@ -1405,6 +1514,7 @@ class RedisCache(BaseCache):
return result
@_redis_circuit_breaker_guard
async def async_lpop(
self,
key: str,
@ -1534,6 +1644,7 @@ class RedisCache(BaseCache):
decoded_results.append(None)
return decoded_results
@_redis_circuit_breaker_guard
async def async_lpop_pipeline(
self,
lpop_list: List[RedisPipelineLpopOperation],

View file

@ -350,6 +350,12 @@ AZURE_DOCUMENT_INTELLIGENCE_DEFAULT_DPI = int(
)
REDIS_SOCKET_TIMEOUT = float(os.getenv("REDIS_SOCKET_TIMEOUT", 0.1))
REDIS_CONNECTION_POOL_TIMEOUT = int(os.getenv("REDIS_CONNECTION_POOL_TIMEOUT", 5))
REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD = int(
os.getenv("REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD", 5)
)
REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT = int(
os.getenv("REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT", 60)
)
# Default Redis major version to assume when version cannot be determined
# Using 7 as it's the modern version that supports LPOP with count parameter
DEFAULT_REDIS_MAJOR_VERSION = int(os.getenv("DEFAULT_REDIS_MAJOR_VERSION", 7))

View file

@ -57,6 +57,22 @@ IMAGE_ATTRIBUTE = "images"
TOOL_CALLS_ATTRIBUTE = "tool_calls"
FUNCTION_CALL_ATTRIBUTE = "function_call"
_SYNC_ITER_EXHAUSTED = object()
def _next_sync_or_exhausted(it: Any) -> Any:
"""
Call next(it) from a thread and return _SYNC_ITER_EXHAUSTED on StopIteration.
asyncio.to_thread re-raises thread exceptions inside a coroutine, where PEP 479
converts StopIteration to RuntimeError before any except clause can catch it.
Returning a sentinel instead keeps StopIteration out of the coroutine boundary.
"""
try:
return next(it)
except StopIteration:
return _SYNC_ITER_EXHAUSTED
def is_async_iterable(obj: Any) -> bool:
"""
@ -2090,7 +2106,9 @@ class CustomStreamWrapper:
):
chunk = self.completion_stream
else:
chunk = next(self.completion_stream) # type: ignore[arg-type]
chunk = await asyncio.to_thread(_next_sync_or_exhausted, self.completion_stream) # type: ignore[arg-type]
if chunk is _SYNC_ITER_EXHAUSTED:
raise StopAsyncIteration
if chunk is not None and chunk != b"":
processed_chunk = self.chunk_creator(chunk=chunk)
if processed_chunk is None:

View file

@ -36,6 +36,8 @@ class MistralOCRConfig(BaseOCRConfig):
- image_min_size: Minimum size of images to include
- bbox_annotation_format: Format for bounding box annotations
- document_annotation_format: Format for document annotations
- extract_header: Whether to extract document header
- extract_footer: Whether to extract document footer
"""
return [
"pages",
@ -44,6 +46,8 @@ class MistralOCRConfig(BaseOCRConfig):
"image_min_size",
"bbox_annotation_format",
"document_annotation_format",
"extract_header",
"extract_footer",
]
def map_ocr_params(

View file

@ -130,6 +130,15 @@ async def get_marketplace():
)
# Allowlist for git-subdir paths: one or more segments separated by '/'.
# Each segment must start with an alphanumeric character and contain only
# alphanumeric characters, dots, hyphens, and underscores.
# This implicitly blocks '..', leading '/', backslashes, and percent-encoded sequences.
_VALID_GIT_SUBDIR_PATH_RE = re.compile(
r"^[a-zA-Z0-9][a-zA-Z0-9._-]*(/[a-zA-Z0-9][a-zA-Z0-9._-]*)*$"
)
@router.post(
"/claude-code/plugins",
tags=["Claude Code Marketplace"],
@ -148,7 +157,7 @@ async def register_plugin(
Parameters:
- name: Plugin name (kebab-case)
- source: Git source reference (github or url format)
- source: Git source reference (github, url, or git-subdir format)
- version: Semantic version (optional)
- description: Plugin description (optional)
- author: Author information (optional)
@ -204,10 +213,34 @@ async def register_plugin(
"error": "URL source must include 'url' field (e.g., 'https://github.com/org/repo.git')"
},
)
elif source_type == "git-subdir":
if not source.get("url"):
raise HTTPException(
status_code=400,
detail={
"error": "git-subdir source must include 'url' field (e.g., 'https://github.com/org/repo.git')"
},
)
if not source.get("path"):
raise HTTPException(
status_code=400,
detail={
"error": "git-subdir source must include 'path' field (e.g., 'plugins/plugin-name')"
},
)
if not _VALID_GIT_SUBDIR_PATH_RE.match(source["path"]):
raise HTTPException(
status_code=400,
detail={
"error": "git-subdir 'path' must be a relative path of the form 'segment/segment' (alphanumeric, dots, hyphens, underscores only)"
},
)
else:
raise HTTPException(
status_code=400,
detail={"error": "source.source must be 'github' or 'url'"},
detail={
"error": "source.source must be 'github', 'url', or 'git-subdir'"
},
)
# Build manifest for storage

View file

@ -920,9 +920,9 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
route=route,
)
if _end_user_object is not None:
end_user_params[
"allowed_model_region"
] = _end_user_object.allowed_model_region
end_user_params["allowed_model_region"] = (
_end_user_object.allowed_model_region
)
if _end_user_object.litellm_budget_table is not None:
_apply_budget_limits_to_end_user_params(
end_user_params=end_user_params,
@ -1416,7 +1416,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
else:
_team_obj = None
user_api_key_cache.set_cache(
await user_api_key_cache.async_set_cache(
key=valid_token.team_id, value=_team_obj
) # save team table in cache - used for tpm/rpm limiting - tpm_rpm_limiter.py
@ -1499,9 +1499,9 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
if _end_user_object is not None:
valid_token_dict.update(end_user_params)
valid_token_dict[
"end_user_object_permission"
] = _end_user_object.object_permission
valid_token_dict["end_user_object_permission"] = (
_end_user_object.object_permission
)
# check if token is from litellm-ui, litellm ui makes keys to allow users to login with sso. These keys can only be used for LiteLLM UI functions
# sso/login, ui/login, /key functions and /user functions

View file

@ -39,7 +39,8 @@ class RegisterPluginRequest(BaseModel):
description=(
"Git source reference. Supported formats:\n"
"- GitHub: {'source': 'github', 'repo': 'org/repo'}\n"
"- Git URL: {'source': 'url', 'url': 'https://github.com/org/repo.git'}"
"- Git URL: {'source': 'url', 'url': 'https://github.com/org/repo.git'}\n"
"- Git Subdir: {'source': 'git-subdir', 'url': 'https://github.com/org/repo.git', 'path': 'plugins/plugin-name'}"
),
)
version: Optional[str] = Field("1.0.0", description="Semantic version")

View file

@ -236,3 +236,49 @@ async def test_get_marketplace(mock_prisma_client):
await mock_prisma_client.db.litellm_claudecodeplugintable.delete(
where={"name": plugin_name}
)
@pytest.mark.asyncio
async def test_register_plugin_git_subdir(mock_prisma_client):
"""Test registering a plugin with git-subdir source type."""
setattr(litellm.proxy.proxy_server, "prisma_client", mock_prisma_client)
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
await litellm.proxy.proxy_server.prisma_client.connect()
plugin_name = f"test-subdir-plugin-{int(time.time())}"
request = RegisterPluginRequest(
name=plugin_name,
source={
"source": "git-subdir",
"url": "https://github.com/test-org/monorepo.git",
"path": "plugins/my-plugin",
},
version="1.0.0",
description="Test git-subdir plugin",
)
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-1234",
user_id="test-user",
)
response = await register_plugin(
request=request,
user_api_key_dict=user_api_key_dict,
)
assert response["status"] == "success"
assert response["action"] == "created"
assert response["plugin"]["name"] == plugin_name
assert response["plugin"]["source"]["source"] == "git-subdir"
assert response["plugin"]["source"]["url"] == "https://github.com/test-org/monorepo.git"
assert response["plugin"]["source"]["path"] == "plugins/my-plugin"
assert response["plugin"]["enabled"] is True
# Cleanup
await mock_prisma_client.db.litellm_claudecodeplugintable.delete(
where={"name": plugin_name}
)

View file

@ -159,3 +159,102 @@ async def test_dual_cache_sync_and_async_set_cache_use_same_ttl():
# Both should use default_in_memory_ttl=60, so their expiry times
# should be within a small tolerance of each other
assert abs(sync_expiry - async_expiry) < 1.0
def test_circuit_breaker_opens_after_threshold():
"""Circuit opens after N consecutive Redis failures."""
from litellm.caching.redis_cache import RedisCircuitBreaker
cb = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60)
for _ in range(3):
cb.record_failure()
assert cb._state == "open"
@pytest.mark.asyncio
async def test_circuit_breaker_open_skips_redis():
"""When circuit is open, the guard decorator raises immediately without calling the method."""
from litellm.caching.redis_cache import (
RedisCircuitBreaker,
_redis_circuit_breaker_guard,
)
class FakeRedis:
def __init__(self):
self._circuit_breaker = RedisCircuitBreaker(
failure_threshold=3, recovery_timeout=60
)
self._circuit_breaker._state = "open"
self._circuit_breaker._opened_at = time.time()
self.call_count = 0
@_redis_circuit_breaker_guard
async def do_thing(self):
self.call_count += 1
return "result"
fr = FakeRedis()
with pytest.raises(Exception, match="circuit breaker is open"):
await fr.do_thing()
assert fr.call_count == 0 # method body never executed
def test_circuit_breaker_closes_on_recovery():
"""After recovery_timeout expires, probe is allowed and success closes the circuit."""
from litellm.caching.redis_cache import RedisCircuitBreaker
cb = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60)
cb._state = "open"
cb._opened_at = time.time() - 9999 # recovery timeout long expired
# is_open() should return False to allow a probe through, and transition to HALF_OPEN
assert cb.is_open() is False
assert cb._state == "half_open"
# Successful probe closes the circuit
cb.record_success()
assert cb._state == "closed"
def test_circuit_breaker_half_open_concurrent_calls_are_fast_failed():
"""
Regression test: only ONE probe gets through when the circuit transitions
OPEN → HALF_OPEN. All concurrent callers that check is_open() while the
state is already HALF_OPEN must be fast-failed (return True), not allowed
through as additional probes.
"""
from litellm.caching.redis_cache import RedisCircuitBreaker
cb = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60)
cb._state = "open"
cb._opened_at = time.time() - 9999 # recovery timeout long expired
# First caller: OPEN + expired → transitions to HALF_OPEN, returns False (probe)
assert cb.is_open() is False
assert cb._state == "half_open"
# All subsequent concurrent callers: HALF_OPEN → fast-fail (return True)
for _ in range(10):
assert cb.is_open() is True, "concurrent callers should be fast-failed in HALF_OPEN"
@pytest.mark.asyncio
async def test_async_increment_cache_returns_none_when_no_in_memory_cache_and_redis_fails():
"""
Regression test: when in_memory_cache is None and Redis fails, async_increment_cache
must return None — not the raw increment delta — to avoid silently miscalculating
rate-limit counters.
"""
dc = DualCache()
dc.in_memory_cache = None # type: ignore[assignment] # constructor always creates InMemoryCache, so null it manually
dc.redis_cache = MagicMock()
dc.redis_cache.async_increment = AsyncMock(side_effect=Exception("redis down"))
result = await dc.async_increment_cache("rpm:model:14-05", 1.0, ttl=60)
assert result is None, (
f"Expected None when in_memory_cache is absent and Redis fails, got {result!r}. "
"Returning the delta (1.0) would silently miscalculate rate-limit counters."
)

View file

@ -1679,3 +1679,150 @@ def test_tool_use_not_dropped_when_finish_reason_already_set(
)
assert tool_calls[0].id == "call_1"
assert tool_calls[0].function.name == "get_weather"
@pytest.mark.asyncio
async def test_custom_stream_wrapper_anext_does_not_block_event_loop_for_sync_iterators(
logging_obj: Logging,
):
"""
Regression test: __anext__ must not call blocking next() on a sync iterator on the
event loop thread. This happens for some provider streams which are sync iterators
but used in async contexts (e.g. boto3-style streaming).
"""
class BlockingIterator:
def __init__(self, chunks, delay_s: float):
self._it = iter(chunks)
self._delay_s = delay_s
def __iter__(self):
return self
def __next__(self):
time.sleep(self._delay_s) # simulate blocking I/O
return next(self._it)
test_chunk = ModelResponseStream(
id="chatcmpl-test",
created=int(time.time()),
model="test-model",
object="chat.completion.chunk",
system_fingerprint=None,
choices=[
StreamingChoices(
finish_reason="stop",
index=0,
delta=Delta(
provider_specific_fields=None,
content="hello",
role="assistant",
function_call=None,
tool_calls=None,
audio=None,
),
logprobs=None,
)
],
provider_specific_fields={},
usage=None,
)
# Delay is intentionally > the wait_for timeout used to detect event loop blocking.
wrapper = CustomStreamWrapper(
completion_stream=BlockingIterator([test_chunk], delay_s=0.3),
model="test-model",
logging_obj=logging_obj,
custom_llm_provider="cached_response",
)
tick_event = asyncio.Event()
async def background_tick():
await asyncio.sleep(0.05)
tick_event.set()
# Run the two coroutines concurrently and measure wall time.
# If __anext__ blocks the event loop, background_tick can't run and the gather
# takes the full 0.3 s delay; if non-blocking both finish within ~0.35 s total.
start = asyncio.get_event_loop().time()
out, _ = await asyncio.gather(
wrapper.__anext__(),
background_tick(),
)
elapsed = asyncio.get_event_loop().time() - start
assert isinstance(out, ModelResponseStream)
# background_tick sleeps 0.05 s; total must finish well under 2 × 0.3 s
assert elapsed < 0.5, f"Event loop was likely blocked (elapsed={elapsed:.2f}s)"
@pytest.mark.asyncio
async def test_custom_stream_wrapper_anext_exhaustion_raises_stop_async_iteration(
logging_obj: Logging,
):
"""
PEP 479 regression: when a sync iterator is exhausted, asyncio.to_thread(next, it)
raises StopIteration inside a coroutine, which Python converts to RuntimeError.
The wrapper must catch StopIteration in the thread and raise StopAsyncIteration
in the coroutine instead, so callers get clean stream termination.
"""
class SingleChunkIterator:
def __init__(self, chunk: ModelResponseStream):
self._chunk = chunk
self._done = False
def __iter__(self):
return self
def __next__(self):
if self._done:
raise StopIteration
self._done = True
return self._chunk
test_chunk = ModelResponseStream(
id="chatcmpl-exhaustion-test",
created=int(time.time()),
model="test-model",
object="chat.completion.chunk",
system_fingerprint=None,
choices=[
StreamingChoices(
finish_reason="stop",
index=0,
delta=Delta(
provider_specific_fields=None,
content="done",
role="assistant",
function_call=None,
tool_calls=None,
audio=None,
),
logprobs=None,
)
],
provider_specific_fields={},
usage=None,
)
wrapper = CustomStreamWrapper(
completion_stream=SingleChunkIterator(test_chunk),
model="test-model",
logging_obj=logging_obj,
custom_llm_provider="cached_response",
)
# Drain the wrapper fully. The wrapper's except-handler calls finish_reason_handler()
# on the first StopAsyncIteration (sent_last_chunk=False→True), then re-raises on the
# next call. What must NOT happen is a RuntimeError from PEP 479 converting
# StopIteration (raised inside the thread) to RuntimeError inside the coroutine.
try:
while True:
await wrapper.__anext__()
except StopAsyncIteration:
pass # expected clean termination
except RuntimeError as e:
pytest.fail(f"PEP 479 regression: StopIteration leaked as RuntimeError: {e}")

View file

@ -0,0 +1,81 @@
"""
Unit tests for MistralOCRConfig transformation.
Tests the supported OCR parameters and their mapping behaviour.
No real API calls are made — all tests are fully mocked/local.
"""
import pytest
from litellm.llms.mistral.ocr.transformation import MistralOCRConfig
@pytest.fixture
def config() -> MistralOCRConfig:
return MistralOCRConfig()
MODEL = "mistral-ocr-latest"
class TestGetSupportedOcrParams:
def test_extract_header_in_supported_params(self, config: MistralOCRConfig) -> None:
"""extract_header must be in the Mistral OCR supported params list."""
supported = config.get_supported_ocr_params(model=MODEL)
assert "extract_header" in supported
def test_extract_footer_in_supported_params(self, config: MistralOCRConfig) -> None:
"""extract_footer must be in the Mistral OCR supported params list."""
supported = config.get_supported_ocr_params(model=MODEL)
assert "extract_footer" in supported
def test_existing_params_still_present(self, config: MistralOCRConfig) -> None:
"""Ensure the previously supported params were not accidentally removed."""
supported = config.get_supported_ocr_params(model=MODEL)
for param in [
"pages",
"include_image_base64",
"image_limit",
"image_min_size",
"bbox_annotation_format",
"document_annotation_format",
]:
assert param in supported, f"Previously supported param '{param}' is missing"
class TestMapOcrParams:
def test_extract_header_passed_through(self, config: MistralOCRConfig) -> None:
"""extract_header=True must survive the map_ocr_params filter."""
result = config.map_ocr_params(
non_default_params={"extract_header": True},
optional_params={},
model=MODEL,
)
assert result == {"extract_header": True}
def test_extract_footer_passed_through(self, config: MistralOCRConfig) -> None:
"""extract_footer=True must survive the map_ocr_params filter."""
result = config.map_ocr_params(
non_default_params={"extract_footer": True},
optional_params={},
model=MODEL,
)
assert result == {"extract_footer": True}
def test_extract_header_and_footer_together(self, config: MistralOCRConfig) -> None:
"""Both params can be passed together and are both forwarded."""
result = config.map_ocr_params(
non_default_params={"extract_header": True, "extract_footer": False},
optional_params={},
model=MODEL,
)
assert result == {"extract_header": True, "extract_footer": False}
def test_unknown_param_is_dropped(self, config: MistralOCRConfig) -> None:
"""Parameters not in the supported list must be silently dropped."""
result = config.map_ocr_params(
non_default_params={"extract_header": True, "unsupported_param": "value"},
optional_params={},
model=MODEL,
)
assert "extract_header" in result
assert "unsupported_param" not in result

View file

@ -0,0 +1,222 @@
"""
Unit tests for claude_code_marketplace.py source validation.
Covers the git-subdir source type added alongside the existing github and url types.
"""
import pytest
from fastapi import HTTPException
from unittest.mock import AsyncMock, MagicMock
import litellm
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.proxy_server import LitellmUserRoles
from litellm.types.proxy.claude_code_endpoints import RegisterPluginRequest
from litellm.proxy.anthropic_endpoints.claude_code_endpoints.claude_code_marketplace import (
register_plugin,
)
def _make_mock_prisma():
"""Stateful prisma mock that supports find_unique, create, and update."""
store: dict = {}
mock_client = MagicMock()
mock_client.proxy_logging_obj = MagicMock()
mock_table = MagicMock()
async def _find_unique(where):
return store.get(where.get("name"))
async def _create(data):
record = MagicMock()
record.id = "test-id"
record.name = data["name"]
record.version = data.get("version")
record.description = data.get("description")
record.manifest_json = data.get("manifest_json", "{}")
record.enabled = data.get("enabled", True)
store[data["name"]] = record
return record
async def _update(where, data):
record = store[where["name"]]
for k, v in data.items():
setattr(record, k, v)
return record
mock_table.find_unique = AsyncMock(side_effect=_find_unique)
mock_table.create = AsyncMock(side_effect=_create)
mock_table.update = AsyncMock(side_effect=_update)
mock_client.db.litellm_claudecodeplugintable = mock_table
return mock_client
_USER = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-1234",
user_id="test-user",
)
_GIT_SUBDIR_SOURCE = {
"source": "git-subdir",
"url": "https://github.com/org/monorepo.git",
"path": "plugins/my-plugin",
}
@pytest.mark.asyncio
async def test_register_plugin_git_subdir_success():
"""git-subdir with both url and path fields registers successfully."""
setattr(litellm.proxy.proxy_server, "prisma_client", _make_mock_prisma())
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
request = RegisterPluginRequest(name="my-monorepo-plugin", source=_GIT_SUBDIR_SOURCE)
response = await register_plugin(request=request, user_api_key_dict=_USER)
assert response["status"] == "success"
assert response["action"] == "created"
assert response["plugin"]["source"]["source"] == "git-subdir"
assert response["plugin"]["source"]["path"] == "plugins/my-plugin"
@pytest.mark.asyncio
async def test_register_plugin_git_subdir_update():
"""Registering the same git-subdir plugin twice returns action=updated."""
setattr(litellm.proxy.proxy_server, "prisma_client", _make_mock_prisma())
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
request = RegisterPluginRequest(
name="my-monorepo-plugin", source=_GIT_SUBDIR_SOURCE, version="1.0.0"
)
await register_plugin(request=request, user_api_key_dict=_USER)
request2 = RegisterPluginRequest(
name="my-monorepo-plugin", source=_GIT_SUBDIR_SOURCE, version="2.0.0"
)
response = await register_plugin(request=request2, user_api_key_dict=_USER)
assert response["status"] == "success"
assert response["action"] == "updated"
@pytest.mark.asyncio
async def test_register_plugin_git_subdir_missing_url():
"""git-subdir without url field raises HTTP 400."""
setattr(litellm.proxy.proxy_server, "prisma_client", _make_mock_prisma())
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
request = RegisterPluginRequest(
name="bad-plugin",
source={"source": "git-subdir", "path": "plugins/my-plugin"},
)
with pytest.raises(HTTPException) as exc_info:
await register_plugin(request=request, user_api_key_dict=_USER)
assert exc_info.value.status_code == 400
assert "url" in exc_info.value.detail["error"]
@pytest.mark.asyncio
async def test_register_plugin_git_subdir_empty_url():
"""git-subdir with empty url raises HTTP 400."""
setattr(litellm.proxy.proxy_server, "prisma_client", _make_mock_prisma())
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
request = RegisterPluginRequest(
name="bad-plugin",
source={"source": "git-subdir", "url": "", "path": "plugins/my-plugin"},
)
with pytest.raises(HTTPException) as exc_info:
await register_plugin(request=request, user_api_key_dict=_USER)
assert exc_info.value.status_code == 400
assert "url" in exc_info.value.detail["error"]
@pytest.mark.asyncio
async def test_register_plugin_git_subdir_missing_path():
"""git-subdir without path field raises HTTP 400."""
setattr(litellm.proxy.proxy_server, "prisma_client", _make_mock_prisma())
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
request = RegisterPluginRequest(
name="bad-plugin",
source={"source": "git-subdir", "url": "https://github.com/org/monorepo.git"},
)
with pytest.raises(HTTPException) as exc_info:
await register_plugin(request=request, user_api_key_dict=_USER)
assert exc_info.value.status_code == 400
assert "path" in exc_info.value.detail["error"]
@pytest.mark.asyncio
async def test_register_plugin_git_subdir_empty_path():
"""git-subdir with empty path raises HTTP 400."""
setattr(litellm.proxy.proxy_server, "prisma_client", _make_mock_prisma())
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
request = RegisterPluginRequest(
name="bad-plugin",
source={"source": "git-subdir", "url": "https://github.com/org/monorepo.git", "path": ""},
)
with pytest.raises(HTTPException) as exc_info:
await register_plugin(request=request, user_api_key_dict=_USER)
assert exc_info.value.status_code == 400
assert "path" in exc_info.value.detail["error"]
@pytest.mark.asyncio
async def test_register_plugin_git_subdir_path_traversal():
"""git-subdir with path traversal segments raises HTTP 400."""
setattr(litellm.proxy.proxy_server, "prisma_client", _make_mock_prisma())
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
for bad_path in [
"../../etc/passwd",
"../secrets",
"/absolute/path",
"plugins\\..\\..\\secrets", # backslash traversal
"plugins/%2e%2e/secrets", # percent-encoded traversal
"plugins/%2E%2E/secrets", # uppercase percent-encoded traversal
"plugins/%252e%252e/secrets", # double-encoded traversal
]:
request = RegisterPluginRequest(
name="bad-plugin",
source={
"source": "git-subdir",
"url": "https://github.com/org/monorepo.git",
"path": bad_path,
},
)
with pytest.raises(HTTPException) as exc_info:
await register_plugin(request=request, user_api_key_dict=_USER)
assert exc_info.value.status_code == 400
assert "relative" in exc_info.value.detail["error"]
@pytest.mark.asyncio
async def test_register_plugin_unknown_source_type():
"""Unknown source type raises HTTP 400 listing all valid types."""
setattr(litellm.proxy.proxy_server, "prisma_client", _make_mock_prisma())
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
request = RegisterPluginRequest(
name="bad-plugin",
source={"source": "ftp", "url": "ftp://example.com/repo"},
)
with pytest.raises(HTTPException) as exc_info:
await register_plugin(request=request, user_api_key_dict=_USER)
assert exc_info.value.status_code == 400
assert "git-subdir" in exc_info.value.detail["error"]

View file

@ -72,7 +72,7 @@ def test_get_api_key_with_custom_litellm_key_header(
def test_team_metadata_with_tags_flows_through_jwt_auth():
"""
Test that team_metadata (specifically tags) flows through JWT authentication.
This is a regression test for the issue where JWT auth was not populating
team_metadata, causing team-level tags to be missing in litellm_pre_call_utils.py
"""
@ -87,7 +87,7 @@ def test_team_metadata_with_tags_flows_through_jwt_auth():
rpm_limit=100,
models=["gpt-4", "gpt-3.5-turbo"],
)
# Simulate constructing UserAPIKeyAuth like we do in JWT auth
# This is the pattern from user_api_key_auth.py lines 552-587
user_api_key_auth = UserAPIKeyAuth(
@ -100,14 +100,16 @@ def test_team_metadata_with_tags_flows_through_jwt_auth():
user_role="internal_user",
user_id="test-user",
)
# Verify team_metadata is set
assert user_api_key_auth.team_metadata is not None, "team_metadata should be populated"
assert (
user_api_key_auth.team_metadata is not None
), "team_metadata should be populated"
assert user_api_key_auth.team_metadata == team_object.metadata, (
f"team_metadata not correctly mapped. "
f"Expected: {team_object.metadata}, Got: {user_api_key_auth.team_metadata}"
)
# Specifically verify tags are present
assert "tags" in user_api_key_auth.team_metadata, "tags should be in team_metadata"
assert user_api_key_auth.team_metadata["tags"] == ["production", "high-priority"], (
@ -118,7 +120,7 @@ def test_team_metadata_with_tags_flows_through_jwt_auth():
def test_route_checks_is_llm_api_route():
"""Test RouteChecks.is_llm_api_route() correctly identifies LLM API routes including passthrough endpoints"""
# Test OpenAI routes
openai_routes = [
"/v1/chat/completions",
@ -142,18 +144,22 @@ def test_route_checks_is_llm_api_route():
"/v1/realtime",
"/realtime",
]
for route in openai_routes:
assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route"
assert RouteChecks.is_llm_api_route(
route=route
), f"Route {route} should be identified as LLM API route"
# Test Anthropic routes
anthropic_routes = [
"/v1/messages",
"/v1/messages/count_tokens",
]
for route in anthropic_routes:
assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route"
assert RouteChecks.is_llm_api_route(
route=route
), f"Route {route} should be identified as LLM API route"
# Test passthrough routes (this is the key improvement over the old route checking)
passthrough_routes = [
@ -171,9 +177,11 @@ def test_route_checks_is_llm_api_route():
"/vllm/v1/chat/completions",
"/mistral/v1/chat/completions",
]
for route in passthrough_routes:
assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route"
assert RouteChecks.is_llm_api_route(
route=route
), f"Route {route} should be identified as LLM API route"
# Test MCP routes
mcp_routes = [
@ -181,9 +189,11 @@ def test_route_checks_is_llm_api_route():
"/mcp/",
"/mcp/test",
]
for route in mcp_routes:
assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route"
assert RouteChecks.is_llm_api_route(
route=route
), f"Route {route} should be identified as LLM API route"
# Test LiteLLM native RAG routes
rag_routes = [
@ -193,7 +203,9 @@ def test_route_checks_is_llm_api_route():
"/v1/rag/query",
]
for route in rag_routes:
assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route"
assert RouteChecks.is_llm_api_route(
route=route
), f"Route {route} should be identified as LLM API route"
# Test routes with placeholders
placeholder_routes = [
@ -206,9 +218,11 @@ def test_route_checks_is_llm_api_route():
"/v1/batches/batch_123",
"/batches/batch_123",
]
for route in placeholder_routes:
assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route"
assert RouteChecks.is_llm_api_route(
route=route
), f"Route {route} should be identified as LLM API route"
# Test Azure OpenAI routes
azure_routes = [
@ -217,9 +231,11 @@ def test_route_checks_is_llm_api_route():
"/engines/gpt-4/chat/completions",
"/engines/gpt-3.5-turbo/completions",
]
for route in azure_routes:
assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route"
assert RouteChecks.is_llm_api_route(
route=route
), f"Route {route} should be identified as LLM API route"
# Test non-LLM routes (should return False)
non_llm_routes = [
@ -236,9 +252,11 @@ def test_route_checks_is_llm_api_route():
"/debug",
"/test",
]
for route in non_llm_routes:
assert not RouteChecks.is_llm_api_route(route=route), f"Route {route} should NOT be identified as LLM API route"
assert not RouteChecks.is_llm_api_route(
route=route
), f"Route {route} should NOT be identified as LLM API route"
# Test invalid inputs
invalid_inputs = [
@ -248,9 +266,11 @@ def test_route_checks_is_llm_api_route():
{},
"",
]
for invalid_input in invalid_inputs:
assert not RouteChecks.is_llm_api_route(route=invalid_input), f"Invalid input {invalid_input} should return False"
assert not RouteChecks.is_llm_api_route(
route=invalid_input
), f"Invalid input {invalid_input} should return False"
@pytest.mark.asyncio
@ -259,7 +279,7 @@ async def test_proxy_admin_expired_key_from_cache():
Test that PROXY_ADMIN keys retrieved from cache are checked for expiration
before being returned. This prevents expired keys from bypassing expiration checks
when retrieved from cache (which normally happens at lines 1014-1036).
Regression test for issue where PROXY_ADMIN keys from cache skipped expiration check.
"""
from datetime import datetime, timedelta, timezone
@ -280,39 +300,43 @@ async def test_proxy_admin_expired_key_from_cache():
api_key = "sk-test-proxy-admin-key"
hashed_key = hash_token(api_key)
expired_time = datetime.now(timezone.utc) - timedelta(hours=1) # Expired 1 hour ago
expired_token = UserAPIKeyAuth(
api_key=api_key,
user_role=LitellmUserRoles.PROXY_ADMIN,
expires=expired_time,
token=hashed_key,
)
# Mock cache to return the expired token
mock_cache = AsyncMock()
mock_cache.async_get_cache = AsyncMock(return_value=expired_token)
mock_cache.delete_cache = MagicMock()
# Mock proxy_logging_obj
mock_proxy_logging_obj = MagicMock()
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = (
AsyncMock()
)
# Mock post_call_failure_hook as async function returning None (no transformation)
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
# Mock prisma_client
mock_prisma_client = MagicMock()
# Mock get_key_object to return expired token from cache
with patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object",
new_callable=AsyncMock,
) as mock_get_key_object, \
patch("litellm.proxy.auth.user_api_key_auth._delete_cache_key_object", new_callable=AsyncMock) as mock_delete_cache:
) as mock_get_key_object, patch(
"litellm.proxy.auth.user_api_key_auth._delete_cache_key_object",
new_callable=AsyncMock,
) as mock_delete_cache:
mock_get_key_object.return_value = expired_token
# Set attributes on proxy_server module (these are imported inside _user_api_key_auth_builder)
import litellm.proxy.proxy_server as _proxy_server_mod
@ -331,14 +355,12 @@ async def test_proxy_admin_expired_key_from_cache():
"litellm_proxy_admin_name": "admin",
}
_original_values = {
attr: getattr(_proxy_server_mod, attr, None)
for attr in _attrs_to_set
attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set
}
try:
for attr, val in _attrs_to_set.items():
setattr(_proxy_server_mod, attr, val)
# Create a mock request
request = Request(scope={"type": "http"})
request._url = URL(url="/chat/completions")
@ -358,38 +380,41 @@ async def test_proxy_admin_expired_key_from_cache():
)
# Verify that ProxyException was raised with expired_key type
assert hasattr(exc_info.value, "type"), "Exception should have 'type' attribute"
assert exc_info.value.type == ProxyErrorTypes.expired_key, (
f"Expected expired_key error type, got {exc_info.value.type}"
)
assert "Expired Key" in str(exc_info.value.message), (
f"Exception message should mention 'Expired Key', got: {exc_info.value.message}"
)
assert hasattr(
exc_info.value, "type"
), "Exception should have 'type' attribute"
assert (
exc_info.value.type == ProxyErrorTypes.expired_key
), f"Expected expired_key error type, got {exc_info.value.type}"
assert "Expired Key" in str(
exc_info.value.message
), f"Exception message should mention 'Expired Key', got: {exc_info.value.message}"
# Verify that the param field does NOT leak the full API key (Issue #18731)
# The param should be abbreviated like "sk-...XXXX" not the full plaintext key
assert exc_info.value.param is not None, "Exception should have 'param' attribute"
assert (
exc_info.value.param is not None
), "Exception should have 'param' attribute"
assert exc_info.value.param != api_key, (
f"SECURITY: Full API key should NOT be in param field! "
f"Got: {exc_info.value.param}, Expected abbreviated format like 'sk-...XXXX'"
)
assert exc_info.value.param.startswith("sk-..."), (
f"Param should be abbreviated to 'sk-...XXXX' format. Got: {exc_info.value.param}"
)
assert exc_info.value.param.startswith(
"sk-..."
), f"Param should be abbreviated to 'sk-...XXXX' format. Got: {exc_info.value.param}"
# Verify that cache deletion was called
mock_delete_cache.assert_called_once()
call_args = mock_delete_cache.call_args
assert call_args[1]["hashed_token"] == hashed_key, (
"Cache deletion should be called with the hashed key"
)
assert (
call_args[1]["hashed_token"] == hashed_key
), "Cache deletion should be called with the hashed key"
finally:
# Restore all module-level attributes so subsequent tests are not affected
for attr, val in _original_values.items():
setattr(_proxy_server_mod, attr, val)
@pytest.mark.asyncio
async def test_return_user_api_key_auth_obj_user_spend_and_budget():
"""
@ -400,7 +425,7 @@ async def test_return_user_api_key_auth_obj_user_spend_and_budget():
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import _return_user_api_key_auth_obj
user_obj = type(
"LiteLLM_UserTable",
(),
@ -413,7 +438,7 @@ async def test_return_user_api_key_auth_obj_user_spend_and_budget():
"user_role": "internal_user",
},
)
api_key = "sk-test-key"
valid_token_dict = {
"user_id": "test-user",
@ -421,10 +446,10 @@ async def test_return_user_api_key_auth_obj_user_spend_and_budget():
}
route = "/chat/completions"
start_time = datetime.now()
mock_service_logger = MagicMock()
mock_service_logger.async_service_success_hook = AsyncMock()
with patch(
"litellm.proxy.auth.user_api_key_auth.user_api_key_service_logger_obj",
new=mock_service_logger,
@ -438,7 +463,7 @@ async def test_return_user_api_key_auth_obj_user_spend_and_budget():
start_time=start_time,
user_role=None,
)
assert isinstance(result, UserAPIKeyAuth)
assert result.user_spend == 250.0
assert result.user_max_budget == 1000.0
@ -470,9 +495,7 @@ def test_proxy_admin_jwt_auth_includes_identity_fields():
user_role=LitellmUserRoles.PROXY_ADMIN,
user_id="user-abc",
team_id="team-123",
team_alias=(
team_object.team_alias if team_object is not None else None
),
team_alias=(team_object.team_alias if team_object is not None else None),
team_metadata=team_object.metadata if team_object is not None else None,
org_id="org-456",
end_user_id="end-user-789",
@ -503,9 +526,7 @@ def test_proxy_admin_jwt_auth_handles_no_team_object():
user_role=LitellmUserRoles.PROXY_ADMIN,
user_id="admin-user",
team_id=None,
team_alias=(
team_object.team_alias if team_object is not None else None
),
team_alias=(team_object.team_alias if team_object is not None else None),
team_metadata=team_object.metadata if team_object is not None else None,
org_id=None,
end_user_id=None,
@ -534,7 +555,10 @@ class TestJWTOAuth2Coexistence:
def test_is_jwt_detects_jwt_tokens(self):
"""JWT tokens have 3 dot-separated parts."""
assert JWTHandler.is_jwt("header.payload.signature") is True
assert JWTHandler.is_jwt("eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.sig123") is True
assert (
JWTHandler.is_jwt("eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.sig123")
is True
)
def test_is_jwt_rejects_opaque_tokens(self):
"""Opaque OAuth2 tokens do not have 3 dot-separated parts."""
@ -567,12 +591,20 @@ class TestJWTOAuth2Coexistence:
mock_request.headers = {"authorization": f"Bearer {opaque_token}"}
mock_request.query_params = {}
with patch("litellm.proxy.proxy_server.general_settings", general_settings), \
patch("litellm.proxy.proxy_server.premium_user", True), \
patch("litellm.proxy.proxy_server.master_key", "sk-master"), \
patch("litellm.proxy.proxy_server.prisma_client", None), \
patch("litellm.proxy.auth.user_api_key_auth.Oauth2Handler.check_oauth2_token", new_callable=AsyncMock, return_value=mock_oauth2_response) as mock_oauth2, \
patch("litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder", new_callable=AsyncMock) as mock_jwt_auth:
with patch(
"litellm.proxy.proxy_server.general_settings", general_settings
), patch("litellm.proxy.proxy_server.premium_user", True), patch(
"litellm.proxy.proxy_server.master_key", "sk-master"
), patch(
"litellm.proxy.proxy_server.prisma_client", None
), patch(
"litellm.proxy.auth.user_api_key_auth.Oauth2Handler.check_oauth2_token",
new_callable=AsyncMock,
return_value=mock_oauth2_response,
) as mock_oauth2, patch(
"litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder",
new_callable=AsyncMock,
) as mock_jwt_auth:
litellm.proxy.proxy_server.jwt_handler.update_environment(
prisma_client=None,
@ -624,12 +656,20 @@ class TestJWTOAuth2Coexistence:
mock_request.headers = {"authorization": f"Bearer {jwt_token}"}
mock_request.query_params = {}
with patch("litellm.proxy.proxy_server.general_settings", general_settings), \
patch("litellm.proxy.proxy_server.premium_user", True), \
patch("litellm.proxy.proxy_server.master_key", "sk-master"), \
patch("litellm.proxy.proxy_server.prisma_client", None), \
patch("litellm.proxy.auth.user_api_key_auth.Oauth2Handler.check_oauth2_token", new_callable=AsyncMock) as mock_oauth2, \
patch("litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder", new_callable=AsyncMock, return_value=mock_jwt_result) as mock_jwt_auth:
with patch(
"litellm.proxy.proxy_server.general_settings", general_settings
), patch("litellm.proxy.proxy_server.premium_user", True), patch(
"litellm.proxy.proxy_server.master_key", "sk-master"
), patch(
"litellm.proxy.proxy_server.prisma_client", None
), patch(
"litellm.proxy.auth.user_api_key_auth.Oauth2Handler.check_oauth2_token",
new_callable=AsyncMock,
) as mock_oauth2, patch(
"litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder",
new_callable=AsyncMock,
return_value=mock_jwt_result,
) as mock_jwt_auth:
litellm.proxy.proxy_server.jwt_handler.update_environment(
prisma_client=None,
@ -671,11 +711,17 @@ class TestJWTOAuth2Coexistence:
mock_request.headers = {"authorization": f"Bearer {jwt_like_token}"}
mock_request.query_params = {}
with patch("litellm.proxy.proxy_server.general_settings", general_settings), \
patch("litellm.proxy.proxy_server.premium_user", True), \
patch("litellm.proxy.proxy_server.master_key", "sk-master"), \
patch("litellm.proxy.proxy_server.prisma_client", None), \
patch("litellm.proxy.auth.user_api_key_auth.Oauth2Handler.check_oauth2_token", new_callable=AsyncMock, return_value=mock_oauth2_response) as mock_oauth2:
with patch(
"litellm.proxy.proxy_server.general_settings", general_settings
), patch("litellm.proxy.proxy_server.premium_user", True), patch(
"litellm.proxy.proxy_server.master_key", "sk-master"
), patch(
"litellm.proxy.proxy_server.prisma_client", None
), patch(
"litellm.proxy.auth.user_api_key_auth.Oauth2Handler.check_oauth2_token",
new_callable=AsyncMock,
return_value=mock_oauth2_response,
) as mock_oauth2:
result = await user_api_key_auth(
request=mock_request,
@ -685,3 +731,131 @@ class TestJWTOAuth2Coexistence:
# OAuth2 should handle it since JWT auth is disabled
mock_oauth2.assert_called_once_with(token=jwt_like_token)
assert result.user_id == "oauth2-user"
@pytest.mark.asyncio
async def test_user_api_key_auth_builder_no_blocking_calls():
"""
_user_api_key_auth_builder must never call any synchronous DualCache method
(set_cache, get_cache, batch_get_cache, increment_cache, delete_cache) on
the hot auth path — those methods call Redis synchronously and block the
event loop. Only async_* variants are allowed.
"""
from starlette.datastructures import URL
from starlette.requests import Request
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder
_blocking_methods = [
"set_cache",
"get_cache",
"batch_get_cache",
"increment_cache",
"delete_cache",
]
api_key = "sk-test-no-blocking-cache"
valid_token = UserAPIKeyAuth(
api_key=api_key,
token=api_key,
user_role=LitellmUserRoles.INTERNAL_USER,
team_id="team-abc",
)
mock_cache = AsyncMock()
mock_cache.async_get_cache = AsyncMock(return_value=valid_token)
mock_cache.async_set_cache = AsyncMock(return_value=None)
# Wire sync methods on the instance as plain MagicMocks (no side_effect) so
# calls are recorded but not raised — the function's broad except Exception
# would swallow a raised error. We assert not_called() after the run instead.
for _m in _blocking_methods:
setattr(mock_cache, _m, MagicMock())
mock_proxy_logging_obj = MagicMock()
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = (
AsyncMock()
)
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
import litellm.proxy.proxy_server as _proxy_server_mod
_attrs = {
"prisma_client": MagicMock(),
"user_api_key_cache": mock_cache,
"proxy_logging_obj": mock_proxy_logging_obj,
"master_key": "sk-master-key",
"general_settings": {},
"llm_model_list": [],
"llm_router": None,
"open_telemetry_logger": None,
"model_max_budget_limiter": MagicMock(),
"user_custom_auth": None,
"jwt_handler": None,
"litellm_proxy_admin_name": "admin",
}
_originals = {k: getattr(_proxy_server_mod, k, None) for k in _attrs}
try:
for k, v in _attrs.items():
setattr(_proxy_server_mod, k, v)
request = Request(scope={"type": "http"})
request._url = URL(url="/chat/completions")
import contextlib
from litellm.caching.dual_cache import DualCache
blocking_patches = [
patch.object(
DualCache,
m,
MagicMock(
side_effect=AssertionError(
f"Blocking DualCache.{m}() called on async hot path — use async_{m}() instead"
)
),
)
for m in _blocking_methods
]
with contextlib.ExitStack() as stack:
for p in blocking_patches:
stack.enter_context(p)
stack.enter_context(
patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object",
new_callable=AsyncMock,
return_value=valid_token,
)
)
stack.enter_context(
patch(
"litellm.proxy.auth.user_api_key_auth.get_team_object",
new_callable=AsyncMock,
return_value=None,
)
)
await _user_api_key_auth_builder(
request=request,
api_key=f"Bearer {api_key}",
azure_api_key_header="",
anthropic_api_key_header=None,
google_ai_studio_api_key_header=None,
azure_apim_header=None,
request_data={},
)
for _m in _blocking_methods:
mock = getattr(mock_cache, _m)
assert mock.call_count == 0, (
f"Blocking DualCache.{_m}() was called {mock.call_count} time(s) "
f"on the async hot path — use async_{_m}() instead"
)
finally:
for k, v in _originals.items():
setattr(_proxy_server_mod, k, v)

View file

@ -0,0 +1,131 @@
import React from "react";
import { act, fireEvent, screen, waitFor } from "@testing-library/react";
import { describe, it, expect, vi, beforeEach } from "vitest";
import { renderWithProviders } from "../../../tests/test-utils";
import AddPluginForm from "./add_plugin_form";
vi.mock("../networking", () => ({
registerClaudeCodePlugin: vi.fn().mockResolvedValue({ status: "success" }),
}));
const DEFAULT_PROPS = {
visible: true,
onClose: vi.fn(),
accessToken: "sk-test",
onSuccess: vi.fn(),
};
describe("AddPluginForm", () => {
beforeEach(() => {
vi.clearAllMocks();
});
it("renders the source type select with GitHub as default", () => {
renderWithProviders(<AddPluginForm {...DEFAULT_PROPS} />);
// The default value "GitHub" is displayed in the collapsed select
expect(screen.getByText("GitHub")).toBeInTheDocument();
// The form label is present
expect(screen.getByText("Source Type")).toBeInTheDocument();
});
it("shows URL and Path fields when git-subdir is selected", async () => {
renderWithProviders(<AddPluginForm {...DEFAULT_PROPS} />);
const sourceSelect = screen.getByLabelText("Source Type");
await act(async () => {
fireEvent.mouseDown(sourceSelect);
});
await waitFor(() => {
fireEvent.click(screen.getByText("Git Subdir"));
});
await waitFor(() => {
expect(screen.getByPlaceholderText("https://github.com/org/repo.git")).toBeInTheDocument();
expect(screen.getByPlaceholderText("plugins/plugin-name")).toBeInTheDocument();
});
});
it("does not show Path field for url source type", async () => {
renderWithProviders(<AddPluginForm {...DEFAULT_PROPS} />);
const sourceSelect = screen.getByLabelText("Source Type");
await act(async () => {
fireEvent.mouseDown(sourceSelect);
});
await waitFor(() => {
fireEvent.click(screen.getByText("Git URL"));
});
await waitFor(() => {
expect(screen.getByPlaceholderText("https://github.com/org/repo.git")).toBeInTheDocument();
expect(screen.queryByPlaceholderText("plugins/plugin-name")).not.toBeInTheDocument();
});
});
it("shows path format error when pattern does not match", async () => {
renderWithProviders(<AddPluginForm {...DEFAULT_PROPS} />);
// Switch to git-subdir
const sourceSelect = screen.getByLabelText("Source Type");
await act(async () => {
fireEvent.mouseDown(sourceSelect);
});
await waitFor(() => {
fireEvent.click(screen.getByText("Git Subdir"));
});
// Fill required fields
fireEvent.change(screen.getByPlaceholderText("my-awesome-plugin"), {
target: { value: "my-plugin" },
});
fireEvent.change(screen.getByPlaceholderText("https://github.com/org/repo.git"), {
target: { value: "https://github.com/org/repo.git" },
});
// Enter a path that violates the allowlist
fireEvent.change(screen.getByPlaceholderText("plugins/plugin-name"), {
target: { value: "../../etc/passwd" },
});
// Submit — triggers Antd form validation
await act(async () => {
fireEvent.click(screen.getByText("Register Plugin"));
});
await waitFor(() => {
expect(
screen.getByText(
"Path must be relative segments (alphanumeric, dots, hyphens, underscores), e.g. plugins/plugin-name"
)
).toBeInTheDocument();
});
});
it("clears path field when switching away from git-subdir", async () => {
renderWithProviders(<AddPluginForm {...DEFAULT_PROPS} />);
// Switch to git-subdir
const sourceSelect = screen.getByLabelText("Source Type");
await act(async () => {
fireEvent.mouseDown(sourceSelect);
});
await waitFor(() => {
fireEvent.click(screen.getByText("Git Subdir"));
});
// Switch back to GitHub
await act(async () => {
fireEvent.mouseDown(sourceSelect);
});
await waitFor(() => {
fireEvent.click(screen.getByText("GitHub"));
});
await waitFor(() => {
expect(screen.queryByPlaceholderText("plugins/plugin-name")).not.toBeInTheDocument();
expect(screen.getByPlaceholderText("anthropics/claude-code")).toBeInTheDocument();
});
});
});

View file

@ -40,7 +40,7 @@ const AddPluginForm: React.FC<AddPluginFormProps> = ({
}) => {
const [form] = Form.useForm();
const [isSubmitting, setIsSubmitting] = useState(false);
const [sourceType, setSourceType] = useState<"github" | "url">("github");
const [sourceType, setSourceType] = useState<"github" | "url" | "git-subdir">("github");
const handleSubmit = async (values: any) => {
if (!accessToken) {
@ -76,6 +76,12 @@ const AddPluginForm: React.FC<AddPluginFormProps> = ({
return;
}
// Validate git URL for url/git-subdir source types
if ((sourceType === "url" || sourceType === "git-subdir") && values.url && !isValidUrl(values.url)) {
MessageManager.error("Invalid git URL format");
return;
}
setIsSubmitting(true);
try {
// Build plugin data
@ -87,6 +93,12 @@ const AddPluginForm: React.FC<AddPluginFormProps> = ({
source: "github",
repo: values.repo.trim(),
}
: sourceType === "git-subdir"
? {
source: "git-subdir",
url: values.url.trim(),
path: values.path.trim(),
}
: {
source: "url",
url: values.url.trim(),
@ -139,10 +151,10 @@ const AddPluginForm: React.FC<AddPluginFormProps> = ({
onClose();
};
const handleSourceTypeChange = (value: "github" | "url") => {
const handleSourceTypeChange = (value: "github" | "url" | "git-subdir") => {
setSourceType(value);
// Clear repo/url fields when switching
form.setFieldsValue({ repo: undefined, url: undefined });
// Clear repo/url/path fields when switching
form.setFieldsValue({ repo: undefined, url: undefined, path: undefined });
};
return (
@ -186,7 +198,8 @@ const AddPluginForm: React.FC<AddPluginFormProps> = ({
>
<Select onChange={handleSourceTypeChange} className="rounded-lg">
<Option value="github">GitHub</Option>
<Option value="url">URL</Option>
<Option value="url">Git URL</Option>
<Option value="git-subdir">Git Subdir</Option>
</Select>
</Form.Item>
@ -209,7 +222,7 @@ const AddPluginForm: React.FC<AddPluginFormProps> = ({
)}
{/* Git URL */}
{sourceType === "url" && (
{(sourceType === "url" || sourceType === "git-subdir") && (
<Form.Item
label="Git URL"
name="url"
@ -224,6 +237,28 @@ const AddPluginForm: React.FC<AddPluginFormProps> = ({
</Form.Item>
)}
{/* Git Subdir Path */}
{sourceType === "git-subdir" && (
<Form.Item
label="Subdirectory Path"
name="path"
rules={[
{ required: true, message: "Please enter subdirectory path" },
{
pattern: /^[a-zA-Z0-9][a-zA-Z0-9._-]*(\/[a-zA-Z0-9][a-zA-Z0-9._-]*)*$/,
message:
"Path must be relative segments (alphanumeric, dots, hyphens, underscores), e.g. plugins/plugin-name",
},
]}
tooltip="Path to the plugin directory within the repository (e.g., plugins/plugin-name)"
>
<Input
placeholder="plugins/plugin-name"
className="rounded-lg"
/>
</Form.Item>
)}
{/* Version */}
<Form.Item
label="Version (Optional)"

View file

@ -15,8 +15,8 @@ interface UseKeyboardNavigationProps {
* Handles J/K for next/previous and Escape for close.
*
* Keyboard shortcuts:
* - J: Navigate to previous log (up)
* - K: Navigate to next log (down)
* - J: Navigate to next log (down)
* - K: Navigate to previous log (up)
* - Escape: Close drawer
*/
export function useKeyboardNavigation({
@ -41,11 +41,11 @@ export function useKeyboardNavigation({
break;
case KEY_J_LOWER:
case KEY_J_UPPER:
selectPreviousLog();
selectNextLog();
break;
case KEY_K_LOWER:
case KEY_K_UPPER:
selectNextLog();
selectPreviousLog();
break;
}
};