diff --git a/docs/my-website/docs/proxy/config_settings.md b/docs/my-website/docs/proxy/config_settings.md
index d7542fc2c3d..16752b030b9 100644
--- a/docs/my-website/docs/proxy/config_settings.md
+++ b/docs/my-website/docs/proxy/config_settings.md
@@ -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
diff --git a/docs/my-website/docs/tutorials/claude_code_plugin_marketplace.md b/docs/my-website/docs/tutorials/claude_code_plugin_marketplace.md
index 9d93c717c4f..d8175f51aca 100644
--- a/docs/my-website/docs/tutorials/claude_code_plugin_marketplace.md
+++ b/docs/my-website/docs/tutorials/claude_code_plugin_marketplace.md
@@ -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.
+
+
+
+```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).
+
diff --git a/litellm/caching/dual_cache.py b/litellm/caching/dual_cache.py
index 4020b8cc22e..34ae3638a5b 100644
--- a/litellm/caching/dual_cache.py
+++ b/litellm/caching/dual_cache.py
@@ -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
diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py
index 82794c116f2..84a2887f527 100644
--- a/litellm/caching/redis_cache.py
+++ b/litellm/caching/redis_cache.py
@@ -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],
diff --git a/litellm/constants.py b/litellm/constants.py
index c0dd115210c..423f01afac1 100644
--- a/litellm/constants.py
+++ b/litellm/constants.py
@@ -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))
diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py
index e99bae8ece8..67e4fadf638 100644
--- a/litellm/litellm_core_utils/streaming_handler.py
+++ b/litellm/litellm_core_utils/streaming_handler.py
@@ -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:
diff --git a/litellm/llms/mistral/ocr/transformation.py b/litellm/llms/mistral/ocr/transformation.py
index 11848f8acf4..3d5e8763027 100644
--- a/litellm/llms/mistral/ocr/transformation.py
+++ b/litellm/llms/mistral/ocr/transformation.py
@@ -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(
diff --git a/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_marketplace.py b/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_marketplace.py
index ab3fa9010e2..abd3ce5661d 100644
--- a/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_marketplace.py
+++ b/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_marketplace.py
@@ -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
diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py
index eb6a5bdb994..e2f06abc52f 100644
--- a/litellm/proxy/auth/user_api_key_auth.py
+++ b/litellm/proxy/auth/user_api_key_auth.py
@@ -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
diff --git a/litellm/types/proxy/claude_code_endpoints.py b/litellm/types/proxy/claude_code_endpoints.py
index 033765527b2..bdf4f122e0e 100644
--- a/litellm/types/proxy/claude_code_endpoints.py
+++ b/litellm/types/proxy/claude_code_endpoints.py
@@ -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")
diff --git a/tests/pass_through_unit_tests/test_claude_code_marketplace.py b/tests/pass_through_unit_tests/test_claude_code_marketplace.py
index 21387439d85..2a1697827b2 100644
--- a/tests/pass_through_unit_tests/test_claude_code_marketplace.py
+++ b/tests/pass_through_unit_tests/test_claude_code_marketplace.py
@@ -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}
+ )
diff --git a/tests/test_litellm/caching/test_dual_cache.py b/tests/test_litellm/caching/test_dual_cache.py
index 606f25ddf44..6bf4307c9cc 100644
--- a/tests/test_litellm/caching/test_dual_cache.py
+++ b/tests/test_litellm/caching/test_dual_cache.py
@@ -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."
+ )
diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py
index e0862629947..20e064ef8f4 100644
--- a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py
+++ b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py
@@ -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}")
diff --git a/tests/test_litellm/llms/mistral/__init__.py b/tests/test_litellm/llms/mistral/__init__.py
new file mode 100644
index 00000000000..e69de29bb2d
diff --git a/tests/test_litellm/llms/mistral/ocr/__init__.py b/tests/test_litellm/llms/mistral/ocr/__init__.py
new file mode 100644
index 00000000000..e69de29bb2d
diff --git a/tests/test_litellm/llms/mistral/ocr/test_mistral_ocr_transformation.py b/tests/test_litellm/llms/mistral/ocr/test_mistral_ocr_transformation.py
new file mode 100644
index 00000000000..ca823d6fb55
--- /dev/null
+++ b/tests/test_litellm/llms/mistral/ocr/test_mistral_ocr_transformation.py
@@ -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
diff --git a/tests/test_litellm/proxy/anthropic_endpoints/test_claude_code_marketplace.py b/tests/test_litellm/proxy/anthropic_endpoints/test_claude_code_marketplace.py
new file mode 100644
index 00000000000..01e0b97138e
--- /dev/null
+++ b/tests/test_litellm/proxy/anthropic_endpoints/test_claude_code_marketplace.py
@@ -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"]
diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py
index f3f0ba56cb9..81ca758983b 100644
--- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py
+++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py
@@ -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)
diff --git a/ui/litellm-dashboard/src/components/claude_code_plugins/add_plugin_form.test.tsx b/ui/litellm-dashboard/src/components/claude_code_plugins/add_plugin_form.test.tsx
new file mode 100644
index 00000000000..36001224cc6
--- /dev/null
+++ b/ui/litellm-dashboard/src/components/claude_code_plugins/add_plugin_form.test.tsx
@@ -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();
+
+ // 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();
+
+ 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();
+
+ 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();
+
+ // 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();
+
+ // 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();
+ });
+ });
+});
diff --git a/ui/litellm-dashboard/src/components/claude_code_plugins/add_plugin_form.tsx b/ui/litellm-dashboard/src/components/claude_code_plugins/add_plugin_form.tsx
index d5e417a9a85..d6a3f0f4fd9 100644
--- a/ui/litellm-dashboard/src/components/claude_code_plugins/add_plugin_form.tsx
+++ b/ui/litellm-dashboard/src/components/claude_code_plugins/add_plugin_form.tsx
@@ -40,7 +40,7 @@ const AddPluginForm: React.FC = ({
}) => {
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 = ({
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 = ({
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 = ({
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 = ({
>
@@ -209,7 +222,7 @@ const AddPluginForm: React.FC = ({
)}
{/* Git URL */}
- {sourceType === "url" && (
+ {(sourceType === "url" || sourceType === "git-subdir") && (
= ({
)}
+ {/* Git Subdir Path */}
+ {sourceType === "git-subdir" && (
+
+
+
+ )}
+
{/* Version */}