From 2ea9e207bd046d5da84a62c54737824e60df5063 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Sat, 21 Mar 2026 12:40:11 -0700 Subject: [PATCH] Litellm ishaan march 20 (#24303) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * 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 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 * [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 * 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 --------- Co-authored-by: Emerson Gomes Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> Co-authored-by: Vincenzo Barrea Co-authored-by: Claude Sonnet 4.6 Co-authored-by: Robert Kirscht Co-authored-by: Imgyu Kim Co-authored-by: Cursor Agent Co-authored-by: Ishaan Jaff --- docs/my-website/docs/proxy/config_settings.md | 2 + .../claude_code_plugin_marketplace.md | 18 +- litellm/caching/dual_cache.py | 21 +- litellm/caching/redis_cache.py | 113 +++++- litellm/constants.py | 6 + .../litellm_core_utils/streaming_handler.py | 20 +- litellm/llms/mistral/ocr/transformation.py | 4 + .../claude_code_marketplace.py | 37 +- litellm/proxy/auth/user_api_key_auth.py | 14 +- litellm/types/proxy/claude_code_endpoints.py | 3 +- .../test_claude_code_marketplace.py | 46 +++ tests/test_litellm/caching/test_dual_cache.py | 99 ++++++ .../test_streaming_handler.py | 147 ++++++++ tests/test_litellm/llms/mistral/__init__.py | 0 .../test_litellm/llms/mistral/ocr/__init__.py | 0 .../ocr/test_mistral_ocr_transformation.py | 81 +++++ .../test_claude_code_marketplace.py | 222 ++++++++++++ .../proxy/auth/test_user_api_key_auth.py | 336 +++++++++++++----- .../add_plugin_form.test.tsx | 131 +++++++ .../claude_code_plugins/add_plugin_form.tsx | 47 ++- .../LogDetailsDrawer/useKeyboardNavigation.ts | 8 +- 21 files changed, 1245 insertions(+), 110 deletions(-) create mode 100644 tests/test_litellm/llms/mistral/__init__.py create mode 100644 tests/test_litellm/llms/mistral/ocr/__init__.py create mode 100644 tests/test_litellm/llms/mistral/ocr/test_mistral_ocr_transformation.py create mode 100644 tests/test_litellm/proxy/anthropic_endpoints/test_claude_code_marketplace.py create mode 100644 ui/litellm-dashboard/src/components/claude_code_plugins/add_plugin_form.test.tsx 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 */}