mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Litellm ishaan march 20 (#24303)
* feat(redis): add circuit breaker to RedisCache to fast-fail when Redis is down (#24181) * feat(redis): add circuit breaker env var constants * feat(redis): add RedisCircuitBreaker and apply guard decorator to all async ops * fix(dual_cache): fall back to L1 instead of re-raising on Redis increment failures * test(caching): add circuit breaker unit tests * fix(redis): fast-fail concurrent HALF_OPEN probes — only one probe at a time * fix(dual_cache): return None fallback when in_memory_cache is absent and Redis fails * test(caching): add regression tests for HALF_OPEN concurrency and None fallback * Fix blocking sync next in __anext__ (#24177) * Fix blocking sync next * Update tests/test_litellm/litellm_core_utils/test_streaming_handler.py Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> * fix PEP 479 regression in __anext__ sync iterator exhaustion asyncio.to_thread re-raises thread exceptions inside a coroutine, where PEP 479 converts StopIteration to RuntimeError before any except clause can catch it. Add _next_sync_or_exhausted() module-level helper that catches StopIteration in the thread and returns a sentinel instead, then raise StopAsyncIteration in the coroutine. Also rewrites the non-blocking test to use asyncio.gather() instead of asyncio.create_task() (which returned None on Python 3.9 / pytest-asyncio in CI), and adds an exhaustion regression test that drains the wrapper fully and asserts no RuntimeError leaks out. --------- Co-authored-by: Emerson Gomes <emerson.gomes@thalesgroup.com> Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> * feat: add git-subdir source type to claude-code/plugins API (#24223) Support a third plugin source type `git-subdir` alongside the existing `github` and `url` types, as documented in the official Claude Code plugin marketplaces spec. New format: {"source": "git-subdir", "url": "...", "path": "subdir/path"} - Validates url and path fields are present and non-empty - Rejects absolute paths, '..' segments, backslashes, and percent-encoded traversal sequences (including double-encoded variants via regex check) - Extracts path validation into _validate_git_subdir_path() helper - Updates Pydantic field description to document all three source types - Adds isValidUrl() check for url/git-subdir source types in the UI form - Adds "Git Subdir" option to the UI form with a required Path field - Adds unit tests covering success, update, missing/empty fields, path traversal variants, and unknown source type Co-authored-by: Claude Sonnet 4.6 <noreply@anthropic.com> * [FEAT] add extract_header and extract_footer to Mistral OCR supported params (#24213) * docs: add git-subdir source type to claude-code plugin marketplace docs (#24289) * fix(ui): swap J/K keyboard navigation in log details drawer (#24279) (#24286) J should navigate down (next) and K should navigate up (previous), matching vim/standard conventions. * fix: use async_set_cache in user_api_key_auth hot path (#24302) * fix: use async_set_cache in auth hot path to avoid blocking event loop * test: assert no blocking set_cache call in _user_api_key_auth_builder * test: broaden blocking call check to all sync DualCache methods * test: fix regression test to actually catch blocking cache calls * fix: ruff lint unused variable + UI build MessageManager error - litellm/caching/redis_cache.py: remove unused variable 'e' in circuit breaker exception handler (F841) - add_plugin_form.tsx: use MessageManager.error() instead of undefined message.error() for git URL validation Co-authored-by: Ishaan Jaff <ishaan-jaff@users.noreply.github.com> * docs: add REDIS_CIRCUIT_BREAKER env vars to config_settings reference Add REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD and REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT to the environment variables reference table so test_env_keys.py passes. Co-authored-by: Ishaan Jaff <ishaan-jaff@users.noreply.github.com> --------- Co-authored-by: Emerson Gomes <emerson.gomes@thalesgroup.com> Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> Co-authored-by: Vincenzo Barrea <manamana88@users.noreply.github.com> Co-authored-by: Claude Sonnet 4.6 <noreply@anthropic.com> Co-authored-by: Robert Kirscht <rkirscht242@gmail.com> Co-authored-by: Imgyu Kim <kimimgo@gmail.com> Co-authored-by: Cursor Agent <cursoragent@cursor.com> Co-authored-by: Ishaan Jaff <ishaan-jaff@users.noreply.github.com>
This commit is contained in:
parent
b64b0d4b9b
commit
2ea9e207bd
21 changed files with 1245 additions and 110 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -37,7 +37,7 @@ Click **+ Add New Plugin** to register a plugin in your marketplace.
|
|||
Enter the plugin information:
|
||||
|
||||
- **Name**: Plugin identifier (kebab-case, e.g., `my-plugin`)
|
||||
- **Source Type**: Choose GitHub or URL
|
||||
- **Source Type**: Choose GitHub, Git URL, or Git Subdir
|
||||
- **Repository/URL**: The git source (e.g., `org/repo` for GitHub)
|
||||
- **Version**: Semantic version (optional)
|
||||
- **Description**: What the plugin does
|
||||
|
|
@ -216,6 +216,22 @@ curl -X DELETE http://localhost:4000/claude-code/plugins/my-plugin \
|
|||
|
||||
Use this format for GitLab, Bitbucket, or self-hosted git repositories.
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="git-subdir" label="Git Subdir">
|
||||
|
||||
```json
|
||||
{
|
||||
"name": "my-plugin",
|
||||
"source": {
|
||||
"source": "git-subdir",
|
||||
"url": "https://github.com/org/repo.git",
|
||||
"path": "plugins/my-plugin"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
Use this format when your plugin lives in a subdirectory of a git repository. The `path` field must be a relative path of slash-separated segments (alphanumeric, dots, hyphens, underscores only).
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
0
tests/test_litellm/llms/mistral/__init__.py
Normal file
0
tests/test_litellm/llms/mistral/__init__.py
Normal file
0
tests/test_litellm/llms/mistral/ocr/__init__.py
Normal file
0
tests/test_litellm/llms/mistral/ocr/__init__.py
Normal file
|
|
@ -0,0 +1,81 @@
|
|||
"""
|
||||
Unit tests for MistralOCRConfig transformation.
|
||||
|
||||
Tests the supported OCR parameters and their mapping behaviour.
|
||||
No real API calls are made — all tests are fully mocked/local.
|
||||
"""
|
||||
import pytest
|
||||
|
||||
from litellm.llms.mistral.ocr.transformation import MistralOCRConfig
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def config() -> MistralOCRConfig:
|
||||
return MistralOCRConfig()
|
||||
|
||||
|
||||
MODEL = "mistral-ocr-latest"
|
||||
|
||||
|
||||
class TestGetSupportedOcrParams:
|
||||
def test_extract_header_in_supported_params(self, config: MistralOCRConfig) -> None:
|
||||
"""extract_header must be in the Mistral OCR supported params list."""
|
||||
supported = config.get_supported_ocr_params(model=MODEL)
|
||||
assert "extract_header" in supported
|
||||
|
||||
def test_extract_footer_in_supported_params(self, config: MistralOCRConfig) -> None:
|
||||
"""extract_footer must be in the Mistral OCR supported params list."""
|
||||
supported = config.get_supported_ocr_params(model=MODEL)
|
||||
assert "extract_footer" in supported
|
||||
|
||||
def test_existing_params_still_present(self, config: MistralOCRConfig) -> None:
|
||||
"""Ensure the previously supported params were not accidentally removed."""
|
||||
supported = config.get_supported_ocr_params(model=MODEL)
|
||||
for param in [
|
||||
"pages",
|
||||
"include_image_base64",
|
||||
"image_limit",
|
||||
"image_min_size",
|
||||
"bbox_annotation_format",
|
||||
"document_annotation_format",
|
||||
]:
|
||||
assert param in supported, f"Previously supported param '{param}' is missing"
|
||||
|
||||
|
||||
class TestMapOcrParams:
|
||||
def test_extract_header_passed_through(self, config: MistralOCRConfig) -> None:
|
||||
"""extract_header=True must survive the map_ocr_params filter."""
|
||||
result = config.map_ocr_params(
|
||||
non_default_params={"extract_header": True},
|
||||
optional_params={},
|
||||
model=MODEL,
|
||||
)
|
||||
assert result == {"extract_header": True}
|
||||
|
||||
def test_extract_footer_passed_through(self, config: MistralOCRConfig) -> None:
|
||||
"""extract_footer=True must survive the map_ocr_params filter."""
|
||||
result = config.map_ocr_params(
|
||||
non_default_params={"extract_footer": True},
|
||||
optional_params={},
|
||||
model=MODEL,
|
||||
)
|
||||
assert result == {"extract_footer": True}
|
||||
|
||||
def test_extract_header_and_footer_together(self, config: MistralOCRConfig) -> None:
|
||||
"""Both params can be passed together and are both forwarded."""
|
||||
result = config.map_ocr_params(
|
||||
non_default_params={"extract_header": True, "extract_footer": False},
|
||||
optional_params={},
|
||||
model=MODEL,
|
||||
)
|
||||
assert result == {"extract_header": True, "extract_footer": False}
|
||||
|
||||
def test_unknown_param_is_dropped(self, config: MistralOCRConfig) -> None:
|
||||
"""Parameters not in the supported list must be silently dropped."""
|
||||
result = config.map_ocr_params(
|
||||
non_default_params={"extract_header": True, "unsupported_param": "value"},
|
||||
optional_params={},
|
||||
model=MODEL,
|
||||
)
|
||||
assert "extract_header" in result
|
||||
assert "unsupported_param" not in result
|
||||
|
|
@ -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"]
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,131 @@
|
|||
import React from "react";
|
||||
import { act, fireEvent, screen, waitFor } from "@testing-library/react";
|
||||
import { describe, it, expect, vi, beforeEach } from "vitest";
|
||||
import { renderWithProviders } from "../../../tests/test-utils";
|
||||
import AddPluginForm from "./add_plugin_form";
|
||||
|
||||
vi.mock("../networking", () => ({
|
||||
registerClaudeCodePlugin: vi.fn().mockResolvedValue({ status: "success" }),
|
||||
}));
|
||||
|
||||
const DEFAULT_PROPS = {
|
||||
visible: true,
|
||||
onClose: vi.fn(),
|
||||
accessToken: "sk-test",
|
||||
onSuccess: vi.fn(),
|
||||
};
|
||||
|
||||
describe("AddPluginForm", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
});
|
||||
|
||||
it("renders the source type select with GitHub as default", () => {
|
||||
renderWithProviders(<AddPluginForm {...DEFAULT_PROPS} />);
|
||||
|
||||
// The default value "GitHub" is displayed in the collapsed select
|
||||
expect(screen.getByText("GitHub")).toBeInTheDocument();
|
||||
// The form label is present
|
||||
expect(screen.getByText("Source Type")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("shows URL and Path fields when git-subdir is selected", async () => {
|
||||
renderWithProviders(<AddPluginForm {...DEFAULT_PROPS} />);
|
||||
|
||||
const sourceSelect = screen.getByLabelText("Source Type");
|
||||
await act(async () => {
|
||||
fireEvent.mouseDown(sourceSelect);
|
||||
});
|
||||
|
||||
await waitFor(() => {
|
||||
fireEvent.click(screen.getByText("Git Subdir"));
|
||||
});
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByPlaceholderText("https://github.com/org/repo.git")).toBeInTheDocument();
|
||||
expect(screen.getByPlaceholderText("plugins/plugin-name")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
it("does not show Path field for url source type", async () => {
|
||||
renderWithProviders(<AddPluginForm {...DEFAULT_PROPS} />);
|
||||
|
||||
const sourceSelect = screen.getByLabelText("Source Type");
|
||||
await act(async () => {
|
||||
fireEvent.mouseDown(sourceSelect);
|
||||
});
|
||||
|
||||
await waitFor(() => {
|
||||
fireEvent.click(screen.getByText("Git URL"));
|
||||
});
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByPlaceholderText("https://github.com/org/repo.git")).toBeInTheDocument();
|
||||
expect(screen.queryByPlaceholderText("plugins/plugin-name")).not.toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
it("shows path format error when pattern does not match", async () => {
|
||||
renderWithProviders(<AddPluginForm {...DEFAULT_PROPS} />);
|
||||
|
||||
// Switch to git-subdir
|
||||
const sourceSelect = screen.getByLabelText("Source Type");
|
||||
await act(async () => {
|
||||
fireEvent.mouseDown(sourceSelect);
|
||||
});
|
||||
await waitFor(() => {
|
||||
fireEvent.click(screen.getByText("Git Subdir"));
|
||||
});
|
||||
|
||||
// Fill required fields
|
||||
fireEvent.change(screen.getByPlaceholderText("my-awesome-plugin"), {
|
||||
target: { value: "my-plugin" },
|
||||
});
|
||||
fireEvent.change(screen.getByPlaceholderText("https://github.com/org/repo.git"), {
|
||||
target: { value: "https://github.com/org/repo.git" },
|
||||
});
|
||||
// Enter a path that violates the allowlist
|
||||
fireEvent.change(screen.getByPlaceholderText("plugins/plugin-name"), {
|
||||
target: { value: "../../etc/passwd" },
|
||||
});
|
||||
|
||||
// Submit — triggers Antd form validation
|
||||
await act(async () => {
|
||||
fireEvent.click(screen.getByText("Register Plugin"));
|
||||
});
|
||||
|
||||
await waitFor(() => {
|
||||
expect(
|
||||
screen.getByText(
|
||||
"Path must be relative segments (alphanumeric, dots, hyphens, underscores), e.g. plugins/plugin-name"
|
||||
)
|
||||
).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
it("clears path field when switching away from git-subdir", async () => {
|
||||
renderWithProviders(<AddPluginForm {...DEFAULT_PROPS} />);
|
||||
|
||||
// Switch to git-subdir
|
||||
const sourceSelect = screen.getByLabelText("Source Type");
|
||||
await act(async () => {
|
||||
fireEvent.mouseDown(sourceSelect);
|
||||
});
|
||||
await waitFor(() => {
|
||||
fireEvent.click(screen.getByText("Git Subdir"));
|
||||
});
|
||||
|
||||
// Switch back to GitHub
|
||||
await act(async () => {
|
||||
fireEvent.mouseDown(sourceSelect);
|
||||
});
|
||||
await waitFor(() => {
|
||||
fireEvent.click(screen.getByText("GitHub"));
|
||||
});
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.queryByPlaceholderText("plugins/plugin-name")).not.toBeInTheDocument();
|
||||
expect(screen.getByPlaceholderText("anthropics/claude-code")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
@ -40,7 +40,7 @@ const AddPluginForm: React.FC<AddPluginFormProps> = ({
|
|||
}) => {
|
||||
const [form] = Form.useForm();
|
||||
const [isSubmitting, setIsSubmitting] = useState(false);
|
||||
const [sourceType, setSourceType] = useState<"github" | "url">("github");
|
||||
const [sourceType, setSourceType] = useState<"github" | "url" | "git-subdir">("github");
|
||||
|
||||
const handleSubmit = async (values: any) => {
|
||||
if (!accessToken) {
|
||||
|
|
@ -76,6 +76,12 @@ const AddPluginForm: React.FC<AddPluginFormProps> = ({
|
|||
return;
|
||||
}
|
||||
|
||||
// Validate git URL for url/git-subdir source types
|
||||
if ((sourceType === "url" || sourceType === "git-subdir") && values.url && !isValidUrl(values.url)) {
|
||||
MessageManager.error("Invalid git URL format");
|
||||
return;
|
||||
}
|
||||
|
||||
setIsSubmitting(true);
|
||||
try {
|
||||
// Build plugin data
|
||||
|
|
@ -87,6 +93,12 @@ const AddPluginForm: React.FC<AddPluginFormProps> = ({
|
|||
source: "github",
|
||||
repo: values.repo.trim(),
|
||||
}
|
||||
: sourceType === "git-subdir"
|
||||
? {
|
||||
source: "git-subdir",
|
||||
url: values.url.trim(),
|
||||
path: values.path.trim(),
|
||||
}
|
||||
: {
|
||||
source: "url",
|
||||
url: values.url.trim(),
|
||||
|
|
@ -139,10 +151,10 @@ const AddPluginForm: React.FC<AddPluginFormProps> = ({
|
|||
onClose();
|
||||
};
|
||||
|
||||
const handleSourceTypeChange = (value: "github" | "url") => {
|
||||
const handleSourceTypeChange = (value: "github" | "url" | "git-subdir") => {
|
||||
setSourceType(value);
|
||||
// Clear repo/url fields when switching
|
||||
form.setFieldsValue({ repo: undefined, url: undefined });
|
||||
// Clear repo/url/path fields when switching
|
||||
form.setFieldsValue({ repo: undefined, url: undefined, path: undefined });
|
||||
};
|
||||
|
||||
return (
|
||||
|
|
@ -186,7 +198,8 @@ const AddPluginForm: React.FC<AddPluginFormProps> = ({
|
|||
>
|
||||
<Select onChange={handleSourceTypeChange} className="rounded-lg">
|
||||
<Option value="github">GitHub</Option>
|
||||
<Option value="url">URL</Option>
|
||||
<Option value="url">Git URL</Option>
|
||||
<Option value="git-subdir">Git Subdir</Option>
|
||||
</Select>
|
||||
</Form.Item>
|
||||
|
||||
|
|
@ -209,7 +222,7 @@ const AddPluginForm: React.FC<AddPluginFormProps> = ({
|
|||
)}
|
||||
|
||||
{/* Git URL */}
|
||||
{sourceType === "url" && (
|
||||
{(sourceType === "url" || sourceType === "git-subdir") && (
|
||||
<Form.Item
|
||||
label="Git URL"
|
||||
name="url"
|
||||
|
|
@ -224,6 +237,28 @@ const AddPluginForm: React.FC<AddPluginFormProps> = ({
|
|||
</Form.Item>
|
||||
)}
|
||||
|
||||
{/* Git Subdir Path */}
|
||||
{sourceType === "git-subdir" && (
|
||||
<Form.Item
|
||||
label="Subdirectory Path"
|
||||
name="path"
|
||||
rules={[
|
||||
{ required: true, message: "Please enter subdirectory path" },
|
||||
{
|
||||
pattern: /^[a-zA-Z0-9][a-zA-Z0-9._-]*(\/[a-zA-Z0-9][a-zA-Z0-9._-]*)*$/,
|
||||
message:
|
||||
"Path must be relative segments (alphanumeric, dots, hyphens, underscores), e.g. plugins/plugin-name",
|
||||
},
|
||||
]}
|
||||
tooltip="Path to the plugin directory within the repository (e.g., plugins/plugin-name)"
|
||||
>
|
||||
<Input
|
||||
placeholder="plugins/plugin-name"
|
||||
className="rounded-lg"
|
||||
/>
|
||||
</Form.Item>
|
||||
)}
|
||||
|
||||
{/* Version */}
|
||||
<Form.Item
|
||||
label="Version (Optional)"
|
||||
|
|
|
|||
|
|
@ -15,8 +15,8 @@ interface UseKeyboardNavigationProps {
|
|||
* Handles J/K for next/previous and Escape for close.
|
||||
*
|
||||
* Keyboard shortcuts:
|
||||
* - J: Navigate to previous log (up)
|
||||
* - K: Navigate to next log (down)
|
||||
* - J: Navigate to next log (down)
|
||||
* - K: Navigate to previous log (up)
|
||||
* - Escape: Close drawer
|
||||
*/
|
||||
export function useKeyboardNavigation({
|
||||
|
|
@ -41,11 +41,11 @@ export function useKeyboardNavigation({
|
|||
break;
|
||||
case KEY_J_LOWER:
|
||||
case KEY_J_UPPER:
|
||||
selectPreviousLog();
|
||||
selectNextLog();
|
||||
break;
|
||||
case KEY_K_LOWER:
|
||||
case KEY_K_UPPER:
|
||||
selectNextLog();
|
||||
selectPreviousLog();
|
||||
break;
|
||||
}
|
||||
};
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue