mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
Fix duplicate async success log emissions for streaming requests (#21355)
* Fix duplicate async success log emission for streaming calls * fix(logging): close remaining async streaming duplicate emission paths * chore: address PR review follow-up comments * Update litellm/litellm_core_utils/streaming_handler.py Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> * [Test] Add unit tests for 5 untested policy components Adds Vitest + RTL test files for policy_table, policy_templates, guardrail_selection_modal, impact_popover, and add_attachment_form. 53 tests total covering rendering, user interactions, API calls, and conditional UI behavior. Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> * feat(key_management): allow @ in key_alias for email-based aliases (#23003) Adds @ to the _KEY_ALIAS_PATTERN allowed character set so that key aliases like user/user@example.com are accepted. Updates tests to cover email-based alias formats. * [Feat[ extends OAuth2 M2M authentication support to info routes (/key/info, /team/info, /user/info, /model/info) (#22713) * added info_route * greptile pt1 * greptile pt2 * greptile pt3 * fix(caching): check REDIS_CLUSTER_NODES env var in Cache and Router class selection (#22790) When Redis Cluster is configured via the REDIS_CLUSTER_NODES environment variable, Cache.__init__() and Router._create_redis_cache() ignored the env var and always created RedisCache instead of RedisClusterCache. This caused the v3 rate limiter's cluster detection (_is_redis_cluster()) to return False, skipping hash-slot key grouping. The resulting CROSSLOT errors were silently caught, falling back to per-instance in-memory counting — breaking RPM/TPM enforcement across multiple proxy instances. Add REDIS_CLUSTER_NODES env var detection to both Cache.__init__() and Router._create_redis_cache(), matching the existing pattern in _redis.py:215-220. When the env var is set and no explicit startup_nodes parameter is provided, parse it and create RedisClusterCache. Fixes #22748 Related to #20836 * address async logging review feedback * fix websocket async logging duplication --------- Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> Co-authored-by: yuneng-jiang <yuneng.jiang@gmail.com> Co-authored-by: Claude Sonnet 4.6 <noreply@anthropic.com> Co-authored-by: hliu-roblox <87345548+hliu-roblox@users.noreply.github.com> Co-authored-by: v0rtex20k <55466324+v0rtex20k@users.noreply.github.com> Co-authored-by: michelligabriele <gabriele.michelli@icloud.com>
This commit is contained in:
parent
3a2cba43dc
commit
c7b1efd685
27 changed files with 1461 additions and 39 deletions
|
|
@ -142,6 +142,7 @@ class A2AStreamingIterator:
|
|||
cache_hit=None,
|
||||
start_time=self.start_time,
|
||||
end_time=end_time,
|
||||
called_from_async=True,
|
||||
)
|
||||
|
||||
verbose_logger.info(
|
||||
|
|
@ -170,4 +171,3 @@ class A2AStreamingIterator:
|
|||
pass
|
||||
|
||||
return result
|
||||
|
||||
|
|
|
|||
|
|
@ -166,6 +166,14 @@ class Cache:
|
|||
None. Cache is set as a litellm param
|
||||
"""
|
||||
if type == LiteLLMCacheType.REDIS:
|
||||
# Check REDIS_CLUSTER_NODES env var if no explicit startup nodes
|
||||
if not redis_startup_nodes:
|
||||
_env_cluster_nodes = litellm.get_secret("REDIS_CLUSTER_NODES")
|
||||
if _env_cluster_nodes is not None and isinstance(
|
||||
_env_cluster_nodes, str
|
||||
):
|
||||
redis_startup_nodes = json.loads(_env_cluster_nodes)
|
||||
|
||||
if redis_startup_nodes:
|
||||
# Only pass GCP parameters if they are provided
|
||||
cluster_kwargs = {
|
||||
|
|
|
|||
|
|
@ -188,6 +188,7 @@ class InteractionsAPIStreamingIterator(BaseInteractionsAPIStreamingIterator):
|
|||
cache_hit=None,
|
||||
start_time=self.start_time,
|
||||
end_time=datetime.now(),
|
||||
called_from_async=True,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -261,4 +262,3 @@ class SyncInteractionsAPIStreamingIterator(BaseInteractionsAPIStreamingIterator)
|
|||
start_time=self.start_time,
|
||||
end_time=datetime.now(),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -217,6 +217,76 @@ _CUSTOM_PRICING_KEYS: frozenset = frozenset(
|
|||
CustomPricingLiteLLMParams.model_fields.keys()
|
||||
)
|
||||
|
||||
# Explicit allowlist of CallTypes that represent async requests.
|
||||
# Keep this list explicit to avoid accidental classification from naming heuristics.
|
||||
_ASYNC_CALL_TYPE_VALUES: frozenset[str] = frozenset(
|
||||
{
|
||||
CallTypes.aembedding.value,
|
||||
CallTypes.acompletion.value,
|
||||
CallTypes.atext_completion.value,
|
||||
CallTypes.aimage_generation.value,
|
||||
CallTypes.aimage_edit.value,
|
||||
CallTypes.amoderation.value,
|
||||
CallTypes.atranscription.value,
|
||||
CallTypes.aspeech.value,
|
||||
CallTypes.arerank.value,
|
||||
CallTypes.asearch.value,
|
||||
CallTypes.arealtime.value,
|
||||
CallTypes.aresponses_websocket.value,
|
||||
CallTypes.acreate_batch.value,
|
||||
CallTypes.aretrieve_batch.value,
|
||||
CallTypes.acancel_batch.value,
|
||||
CallTypes.anthropic_messages.value,
|
||||
CallTypes.aget_assistants.value,
|
||||
CallTypes.acreate_assistants.value,
|
||||
CallTypes.adelete_assistant.value,
|
||||
CallTypes.acreate_thread.value,
|
||||
CallTypes.aget_thread.value,
|
||||
CallTypes.a_add_message.value,
|
||||
CallTypes.aget_messages.value,
|
||||
CallTypes.arun_thread.value,
|
||||
CallTypes.arun_thread_stream.value,
|
||||
CallTypes.afile_retrieve.value,
|
||||
CallTypes.afile_delete.value,
|
||||
CallTypes.afile_list.value,
|
||||
CallTypes.acreate_file.value,
|
||||
CallTypes.afile_content.value,
|
||||
CallTypes.acreate_fine_tuning_job.value,
|
||||
CallTypes.acreate_video.value,
|
||||
CallTypes.avideo_retrieve.value,
|
||||
CallTypes.avideo_content.value,
|
||||
CallTypes.avideo_remix.value,
|
||||
CallTypes.avideo_list.value,
|
||||
CallTypes.avideo_retrieve_job.value,
|
||||
CallTypes.avideo_delete.value,
|
||||
CallTypes.avector_store_file_create.value,
|
||||
CallTypes.avector_store_file_list.value,
|
||||
CallTypes.avector_store_file_retrieve.value,
|
||||
CallTypes.avector_store_file_content.value,
|
||||
CallTypes.avector_store_file_update.value,
|
||||
CallTypes.avector_store_file_delete.value,
|
||||
CallTypes.avector_store_create.value,
|
||||
CallTypes.avector_store_search.value,
|
||||
CallTypes.acreate_container.value,
|
||||
CallTypes.alist_containers.value,
|
||||
CallTypes.aretrieve_container.value,
|
||||
CallTypes.adelete_container.value,
|
||||
CallTypes.alist_container_files.value,
|
||||
CallTypes.aupload_container_file.value,
|
||||
CallTypes.acancel_fine_tuning_job.value,
|
||||
CallTypes.alist_fine_tuning_jobs.value,
|
||||
CallTypes.aretrieve_fine_tuning_job.value,
|
||||
CallTypes.aresponses.value,
|
||||
CallTypes.alist_input_items.value,
|
||||
CallTypes.allm_passthrough_route.value,
|
||||
CallTypes.agenerate_content.value,
|
||||
CallTypes.agenerate_content_stream.value,
|
||||
CallTypes.aocr.value,
|
||||
CallTypes.asend_message.value,
|
||||
CallTypes.acreate_skill.value,
|
||||
}
|
||||
)
|
||||
|
||||
sentry_sdk_instance = None
|
||||
capture_exception = None
|
||||
add_breadcrumb = None
|
||||
|
|
@ -1581,6 +1651,30 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
|
||||
return True
|
||||
|
||||
def _is_async_litellm_request(self, litellm_params: Dict[str, Any]) -> bool:
|
||||
"""
|
||||
Best-effort async request detection for logging flows.
|
||||
|
||||
We need this to prevent sync handler payload/callback duplication when the same
|
||||
request is processed by both async and sync callback paths.
|
||||
"""
|
||||
if litellm_params.get("async_call", False) is True:
|
||||
return True
|
||||
|
||||
for async_call_type in _ASYNC_CALL_TYPE_VALUES:
|
||||
if litellm_params.get(async_call_type, False) is True:
|
||||
return True
|
||||
|
||||
call_type = self.call_type
|
||||
if isinstance(call_type, CallTypes):
|
||||
call_type_value = call_type.value
|
||||
else:
|
||||
call_type_value = str(call_type)
|
||||
if call_type_value.startswith("CallTypes."):
|
||||
call_type_value = call_type_value.split("CallTypes.", 1)[1]
|
||||
|
||||
return call_type_value in _ASYNC_CALL_TYPE_VALUES
|
||||
|
||||
def _update_completion_start_time(self, completion_start_time: datetime.datetime):
|
||||
self.completion_start_time = completion_start_time
|
||||
self.model_call_details["completion_start_time"] = self.completion_start_time
|
||||
|
|
@ -1918,12 +2012,8 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
standard_logging_object=kwargs.get("standard_logging_object", None),
|
||||
)
|
||||
litellm_params = self.model_call_details.get("litellm_params", {})
|
||||
is_sync_request = (
|
||||
litellm_params.get(CallTypes.acompletion.value, False) is not True
|
||||
and litellm_params.get(CallTypes.aresponses.value, False) is not True
|
||||
and litellm_params.get(CallTypes.aembedding.value, False) is not True
|
||||
and litellm_params.get(CallTypes.aimage_generation.value, False) is not True
|
||||
and litellm_params.get(CallTypes.atranscription.value, False) is not True
|
||||
is_sync_request = not self._is_async_litellm_request(
|
||||
litellm_params=litellm_params
|
||||
)
|
||||
try:
|
||||
## BUILD COMPLETE STREAMED RESPONSE
|
||||
|
|
@ -1961,7 +2051,10 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
)
|
||||
) is not None:
|
||||
# Only emit for sync requests (async_success_handler handles async)
|
||||
if is_sync_request:
|
||||
if (
|
||||
is_sync_request
|
||||
and kwargs.get("called_from_async", False) is not True
|
||||
):
|
||||
emit_standard_logging_payload(standard_logging_payload)
|
||||
callbacks = self.get_combined_callback_list(
|
||||
dynamic_success_callbacks=self.dynamic_success_callbacks,
|
||||
|
|
@ -2821,12 +2914,8 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
): # prevent double logging
|
||||
return
|
||||
litellm_params = self.model_call_details.get("litellm_params", {})
|
||||
is_sync_request = (
|
||||
litellm_params.get(CallTypes.acompletion.value, False) is not True
|
||||
and litellm_params.get(CallTypes.aresponses.value, False) is not True
|
||||
and litellm_params.get(CallTypes.aembedding.value, False) is not True
|
||||
and litellm_params.get(CallTypes.aimage_generation.value, False) is not True
|
||||
and litellm_params.get(CallTypes.atranscription.value, False) is not True
|
||||
is_sync_request = not self._is_async_litellm_request(
|
||||
litellm_params=litellm_params
|
||||
)
|
||||
|
||||
try:
|
||||
|
|
@ -3153,6 +3242,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
start_time,
|
||||
end_time,
|
||||
cache_hit,
|
||||
called_from_async=True,
|
||||
)
|
||||
|
||||
def _should_run_sync_callbacks_for_async_calls(self) -> bool:
|
||||
|
|
|
|||
|
|
@ -218,7 +218,11 @@ class RealTimeStreaming:
|
|||
# Create an event loop for the new thread
|
||||
asyncio.create_task(self.logging_obj.async_success_handler(self.messages))
|
||||
## SYNC LOGGING
|
||||
executor.submit(self.logging_obj.success_handler(self.messages))
|
||||
executor.submit(
|
||||
self.logging_obj.success_handler,
|
||||
self.messages,
|
||||
called_from_async=True,
|
||||
)
|
||||
|
||||
async def _send_to_backend(self, message: str) -> None:
|
||||
"""Send a message to the backend WebSocket.
|
||||
|
|
|
|||
|
|
@ -2042,13 +2042,23 @@ class CustomStreamWrapper:
|
|||
raise StopAsyncIteration
|
||||
else: # temporary patch for non-aiohttp async calls
|
||||
# example - boto3 bedrock llms
|
||||
exhausted_sentinel = object()
|
||||
while True:
|
||||
if isinstance(self.completion_stream, str) or isinstance(
|
||||
self.completion_stream, bytes
|
||||
):
|
||||
chunk = self.completion_stream
|
||||
else:
|
||||
chunk = next(self.completion_stream) # type: ignore[arg-type]
|
||||
# Sync iterators can block (e.g. boto3 streams). Run next()
|
||||
# off the event loop thread to keep async callers responsive.
|
||||
# NOTE: We use a sentinel default because StopIteration cannot
|
||||
# be stored in an asyncio.Future (raises TypeError), which
|
||||
# would cause the await to hang forever.
|
||||
chunk = await asyncio.to_thread( # type: ignore[arg-type]
|
||||
next, self.completion_stream, exhausted_sentinel
|
||||
)
|
||||
if chunk is exhausted_sentinel:
|
||||
raise StopIteration
|
||||
if chunk is not None and chunk != b"":
|
||||
processed_chunk = self.chunk_creator(chunk=chunk)
|
||||
if processed_chunk is None:
|
||||
|
|
@ -2123,6 +2133,7 @@ class CustomStreamWrapper:
|
|||
cache_hit=cache_hit,
|
||||
start_time=None,
|
||||
end_time=None,
|
||||
called_from_async=True,
|
||||
)
|
||||
|
||||
raise StopAsyncIteration # Re-raise StopIteration
|
||||
|
|
|
|||
|
|
@ -839,7 +839,7 @@ def _check_end_user_budget(
|
|||
Raises:
|
||||
litellm.BudgetExceededError: If end user has exceeded their budget
|
||||
"""
|
||||
if route in LiteLLMRoutes.info_routes.value:
|
||||
if RouteChecks.is_info_route(route):
|
||||
return
|
||||
|
||||
if end_user_obj.litellm_budget_table is None:
|
||||
|
|
|
|||
|
|
@ -164,9 +164,8 @@ class RouteChecks:
|
|||
|
||||
if RouteChecks.is_llm_api_route(route=route):
|
||||
pass
|
||||
elif (
|
||||
route in LiteLLMRoutes.info_routes.value
|
||||
): # check if user allowed to call an info route
|
||||
elif RouteChecks.is_info_route(route=route):
|
||||
# check if user allowed to call an info route
|
||||
if route == "/key/info":
|
||||
# handled by function itself
|
||||
pass
|
||||
|
|
@ -358,6 +357,13 @@ class RouteChecks:
|
|||
"""
|
||||
return route in LiteLLMRoutes.management_routes.value
|
||||
|
||||
@staticmethod
|
||||
def is_info_route(route: str) -> bool:
|
||||
"""
|
||||
Check if route is an info route
|
||||
"""
|
||||
return route in LiteLLMRoutes.info_routes.value
|
||||
|
||||
@staticmethod
|
||||
def _is_azure_openai_route(route: str) -> bool:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -632,9 +632,10 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
########## End of Route Checks Before Reading DB / Cache for "token" ########
|
||||
|
||||
if general_settings.get("enable_oauth2_auth", False) is True:
|
||||
# Only apply OAuth2 M2M authentication to LLM API routes, not UI/management routes
|
||||
# Only apply OAuth2 M2M authentication to LLM API routes and info routes, not UI/management routes
|
||||
# This allows UI SSO to work separately from API M2M authentication
|
||||
if RouteChecks.is_llm_api_route(route=route):
|
||||
# Note: Info routes are already scoped to the user
|
||||
if RouteChecks.is_llm_api_route(route=route) or RouteChecks.is_info_route(route=route):
|
||||
# return UserAPIKeyAuth object
|
||||
# helper to check if the api_key is a valid oauth2 token
|
||||
from litellm.proxy.proxy_server import premium_user
|
||||
|
|
|
|||
|
|
@ -38,8 +38,9 @@ class _ProxyDBLogger(CustomLogger):
|
|||
request_route = user_api_key_dict.request_route
|
||||
if _ProxyDBLogger._should_track_errors_in_db() is False:
|
||||
return
|
||||
elif request_route is not None and not RouteChecks.is_llm_api_route(
|
||||
route=request_route
|
||||
elif request_route is not None and not (
|
||||
RouteChecks.is_llm_api_route(route=request_route) or
|
||||
RouteChecks.is_info_route(route=request_route)
|
||||
):
|
||||
return
|
||||
|
||||
|
|
|
|||
|
|
@ -4991,7 +4991,7 @@ async def test_key_logging(
|
|||
)
|
||||
|
||||
|
||||
_KEY_ALIAS_PATTERN = re.compile(r"^[a-zA-Z0-9][a-zA-Z0-9_\-/\.]{0,253}[a-zA-Z0-9]$")
|
||||
_KEY_ALIAS_PATTERN = re.compile(r"^[a-zA-Z0-9][a-zA-Z0-9_\-/\.@]{0,253}[a-zA-Z0-9]$")
|
||||
|
||||
|
||||
def _validate_key_alias_format(key_alias: Optional[str]) -> None:
|
||||
|
|
@ -5009,7 +5009,7 @@ def _validate_key_alias_format(key_alias: Optional[str]) -> None:
|
|||
|
||||
if not _KEY_ALIAS_PATTERN.match(key_alias):
|
||||
raise ProxyException(
|
||||
message="Invalid key_alias format. Must be 2-255 characters, start/end with alphanumeric, and only contain a-zA-Z0-9_-/.",
|
||||
message="Invalid key_alias format. Must be 2-255 characters, start/end with alphanumeric, and only contain a-zA-Z0-9_-/.@.",
|
||||
type=ProxyErrorTypes.bad_request_error,
|
||||
param="key_alias",
|
||||
code=400,
|
||||
|
|
|
|||
|
|
@ -1773,7 +1773,7 @@ class ProxyLogging:
|
|||
"""
|
||||
|
||||
#########################################################
|
||||
# Only log LLM API errors for proxy level hooks
|
||||
# Only log LLM API and info route errors for proxy level hooks
|
||||
# eg. Authentication errors, rate limit errors, etc.
|
||||
# Note: This fixes a security issue where we
|
||||
# would log temporary keys/auth info
|
||||
|
|
@ -1781,7 +1781,10 @@ class ProxyLogging:
|
|||
#########################################################
|
||||
if route is None:
|
||||
return False
|
||||
if RouteChecks.is_llm_api_route(route) is not True:
|
||||
if not (
|
||||
RouteChecks.is_llm_api_route(route) or
|
||||
RouteChecks.is_info_route(route)
|
||||
):
|
||||
return False
|
||||
|
||||
return isinstance(original_exception, HTTPException) or (
|
||||
|
|
|
|||
|
|
@ -468,6 +468,7 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
cache_hit=None,
|
||||
start_time=self.start_time,
|
||||
end_time=datetime.now(),
|
||||
called_from_async=True,
|
||||
)
|
||||
self._run_post_success_hooks(end_time=datetime.now())
|
||||
|
||||
|
|
@ -803,7 +804,11 @@ class ResponsesWebSocketStreaming:
|
|||
asyncio.create_task(
|
||||
self.logging_obj.async_success_handler(self.messages)
|
||||
)
|
||||
_ws_executor.submit(self.logging_obj.success_handler, self.messages)
|
||||
_ws_executor.submit(
|
||||
self.logging_obj.success_handler,
|
||||
self.messages,
|
||||
called_from_async=True,
|
||||
)
|
||||
|
||||
async def backend_to_client(self) -> None:
|
||||
"""Forward events from backend WebSocket to the client."""
|
||||
|
|
|
|||
|
|
@ -728,8 +728,18 @@ class Router:
|
|||
"""
|
||||
Initializes either a RedisCache or RedisClusterCache based on the cache_config.
|
||||
"""
|
||||
if cache_config.get("startup_nodes"):
|
||||
return RedisClusterCache(**cache_config)
|
||||
startup_nodes = cache_config.get("startup_nodes")
|
||||
if not startup_nodes:
|
||||
_env_cluster_nodes = get_secret("REDIS_CLUSTER_NODES")
|
||||
if _env_cluster_nodes is not None and isinstance(
|
||||
_env_cluster_nodes, str
|
||||
):
|
||||
startup_nodes = json.loads(_env_cluster_nodes)
|
||||
|
||||
if startup_nodes:
|
||||
return RedisClusterCache(
|
||||
**{**cache_config, "startup_nodes": startup_nodes}
|
||||
)
|
||||
else:
|
||||
return RedisCache(**cache_config)
|
||||
|
||||
|
|
|
|||
|
|
@ -1796,6 +1796,9 @@ def client(original_function): # noqa: PLR0915
|
|||
assert (
|
||||
logging_obj is not None
|
||||
), "logging_obj should not be None after function_setup"
|
||||
# Mark this logging object as async-driven so sync callback paths can
|
||||
# safely avoid duplicate standard payload/callback emission.
|
||||
logging_obj.litellm_params["async_call"] = True
|
||||
|
||||
modified_kwargs = await async_pre_call_deployment_hook(kwargs, call_type)
|
||||
if modified_kwargs is not None:
|
||||
|
|
@ -1940,11 +1943,6 @@ def client(original_function): # noqa: PLR0915
|
|||
is_completion_with_fallbacks=is_completion_with_fallbacks,
|
||||
)
|
||||
)
|
||||
logging_obj.handle_sync_success_callbacks_for_async_calls(
|
||||
result=result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
# REBUILD EMBEDDING CACHING
|
||||
if (
|
||||
isinstance(result, EmbeddingResponse)
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ sys.path.insert(
|
|||
0, os.path.abspath("../../..")
|
||||
) # Adds the parent directory to the system path
|
||||
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
from litellm.caching.redis_cluster_cache import RedisClusterCache
|
||||
|
||||
|
||||
|
|
@ -64,3 +65,117 @@ async def test_redis_cluster_async_batch_get(mock_init_redis_cluster):
|
|||
# Verify mget_nonatomic was called instead of mget
|
||||
mock_redis.mget_nonatomic.assert_called_once()
|
||||
assert not mock_redis.mget.called
|
||||
|
||||
|
||||
@patch("litellm._redis.get_redis_connection_pool")
|
||||
@patch("litellm._redis.get_redis_client")
|
||||
@patch("litellm.caching.redis_cache.RedisCache._setup_health_pings")
|
||||
def test_cache_init_creates_cluster_cache_from_env_var(
|
||||
mock_health, mock_get_client, mock_get_pool, monkeypatch
|
||||
):
|
||||
"""
|
||||
Test that Cache() creates RedisClusterCache when REDIS_CLUSTER_NODES env var is set.
|
||||
|
||||
Regression test for https://github.com/BerriAI/litellm/issues/22748
|
||||
"""
|
||||
from litellm.caching.caching import Cache
|
||||
|
||||
startup_nodes = [{"host": "127.0.0.1", "port": "7001"}]
|
||||
monkeypatch.setenv("REDIS_CLUSTER_NODES", json.dumps(startup_nodes))
|
||||
monkeypatch.delenv("REDIS_HOST", raising=False)
|
||||
monkeypatch.delenv("REDIS_PORT", raising=False)
|
||||
monkeypatch.delenv("REDIS_PASSWORD", raising=False)
|
||||
monkeypatch.delenv("REDIS_URL", raising=False)
|
||||
|
||||
mock_get_client.return_value = MagicMock()
|
||||
mock_get_pool.return_value = MagicMock()
|
||||
|
||||
cache = Cache(type="redis")
|
||||
assert isinstance(cache.cache, RedisClusterCache)
|
||||
|
||||
|
||||
@patch("litellm._redis.get_redis_connection_pool")
|
||||
@patch("litellm._redis.get_redis_client")
|
||||
@patch("litellm.caching.redis_cache.RedisCache._setup_health_pings")
|
||||
def test_cache_init_creates_redis_cache_without_cluster_config(
|
||||
mock_health, mock_get_client, mock_get_pool, monkeypatch
|
||||
):
|
||||
"""
|
||||
Test that Cache() creates RedisCache when no cluster config is present.
|
||||
|
||||
Ensures backward compatibility: without REDIS_CLUSTER_NODES or
|
||||
redis_startup_nodes, the standard RedisCache is still used.
|
||||
"""
|
||||
from litellm.caching.caching import Cache
|
||||
|
||||
monkeypatch.delenv("REDIS_CLUSTER_NODES", raising=False)
|
||||
monkeypatch.setenv("REDIS_HOST", "localhost")
|
||||
monkeypatch.setenv("REDIS_PORT", "6379")
|
||||
monkeypatch.delenv("REDIS_URL", raising=False)
|
||||
|
||||
mock_get_client.return_value = MagicMock()
|
||||
mock_get_pool.return_value = MagicMock()
|
||||
|
||||
cache = Cache(type="redis")
|
||||
assert isinstance(cache.cache, RedisCache)
|
||||
assert not isinstance(cache.cache, RedisClusterCache)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"startup_nodes, env_var, expected_cache_type",
|
||||
[
|
||||
pytest.param(
|
||||
[dict(host="node1.localhost", port=6379)],
|
||||
None,
|
||||
RedisClusterCache,
|
||||
id="cluster-via-explicit-startup-nodes",
|
||||
),
|
||||
pytest.param(
|
||||
None,
|
||||
'[{"host": "node1.localhost", "port": 6379}]',
|
||||
RedisClusterCache,
|
||||
id="cluster-via-env-var",
|
||||
),
|
||||
pytest.param(
|
||||
None,
|
||||
None,
|
||||
RedisCache,
|
||||
id="standard-redis-when-no-cluster-config",
|
||||
),
|
||||
pytest.param(
|
||||
[dict(host="explicit-node.localhost", port=6379)],
|
||||
'[{"host": "env-node.localhost", "port": 6379}]',
|
||||
RedisClusterCache,
|
||||
id="explicit-startup-nodes-takes-precedence-over-env-var",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_router_create_redis_cache_cluster_detection(
|
||||
startup_nodes, env_var, expected_cache_type, monkeypatch
|
||||
):
|
||||
"""
|
||||
Test that Router._create_redis_cache() creates RedisClusterCache when
|
||||
either startup_nodes is in config or REDIS_CLUSTER_NODES env var is set.
|
||||
Also verifies that explicit startup_nodes take precedence over env var.
|
||||
|
||||
Regression test for https://github.com/BerriAI/litellm/issues/22748
|
||||
"""
|
||||
from litellm import Router
|
||||
|
||||
cache_config = dict(
|
||||
host="mockhost",
|
||||
port=6379,
|
||||
password="mock-password",
|
||||
startup_nodes=startup_nodes,
|
||||
)
|
||||
|
||||
if env_var is not None:
|
||||
monkeypatch.setenv("REDIS_CLUSTER_NODES", env_var)
|
||||
else:
|
||||
monkeypatch.delenv("REDIS_CLUSTER_NODES", raising=False)
|
||||
|
||||
def _mock_redis_cache_init(*args, **kwargs): ...
|
||||
|
||||
with patch.object(RedisCache, "__init__", _mock_redis_cache_init):
|
||||
redis_cache = Router._create_redis_cache(cache_config)
|
||||
assert isinstance(redis_cache, expected_cache_type)
|
||||
|
|
|
|||
|
|
@ -365,6 +365,143 @@ def test_success_handler_skips_sync_callbacks_for_async_requests(logging_obj, as
|
|||
dummy_logger.log_stream_event.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"call_type, expected_async",
|
||||
[
|
||||
("agenerate_content_stream", True),
|
||||
("anthropic_messages", True),
|
||||
("_aresponses_websocket", True),
|
||||
("add_message", False),
|
||||
("completion", False),
|
||||
],
|
||||
)
|
||||
def test_is_async_litellm_request_detection_uses_call_type(logging_obj, call_type, expected_async):
|
||||
logging_obj.call_type = call_type
|
||||
logging_obj.model_call_details["litellm_params"] = {}
|
||||
logging_obj.litellm_params = {}
|
||||
|
||||
assert (
|
||||
logging_obj._is_async_litellm_request(logging_obj.model_call_details["litellm_params"])
|
||||
is expected_async
|
||||
)
|
||||
|
||||
|
||||
def test_success_handler_does_not_emit_standard_payload_for_async_call_marker(logging_obj):
|
||||
"""
|
||||
Regression test for async call types that don't set legacy `litellm_params` flags
|
||||
(e.g. `agenerate_content_stream`).
|
||||
"""
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
logging_obj.stream = True
|
||||
logging_obj.call_type = CallTypes.agenerate_content_stream.value
|
||||
logging_obj.model_call_details["litellm_params"] = {"async_call": True}
|
||||
logging_obj.litellm_params = logging_obj.model_call_details["litellm_params"]
|
||||
|
||||
model_response = ModelResponse(
|
||||
id="resp-123",
|
||||
model="gemini-2.5-pro",
|
||||
choices=[
|
||||
{
|
||||
"message": {"role": "assistant", "content": "hello"},
|
||||
"finish_reason": "stop",
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
usage={"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
logging_obj,
|
||||
"_get_assembled_streaming_response",
|
||||
return_value=model_response,
|
||||
),
|
||||
patch(
|
||||
"litellm.litellm_core_utils.litellm_logging.get_standard_logging_object_payload",
|
||||
return_value={"id": "test-payload"},
|
||||
),
|
||||
patch.object(logging_obj, "get_combined_callback_list", return_value=[]),
|
||||
patch(
|
||||
"litellm.litellm_core_utils.litellm_logging.emit_standard_logging_payload"
|
||||
) as mock_emit,
|
||||
):
|
||||
logging_obj.success_handler(result=model_response)
|
||||
|
||||
mock_emit.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_wrapper_async_calls_sync_success_callbacks_once():
|
||||
"""
|
||||
Regression test: wrapper_async should trigger sync-success callbacks once.
|
||||
|
||||
Historically, wrapper_async invoked handle_sync_success_callbacks_for_async_calls()
|
||||
directly and via _client_async_logging_helper, causing duplicate calls.
|
||||
"""
|
||||
import asyncio
|
||||
|
||||
import litellm
|
||||
|
||||
test_logging_obj = LitellmLogging(
|
||||
model="openai/codex-mini-latest",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
stream=False,
|
||||
call_type="acompletion",
|
||||
start_time=time.time(),
|
||||
litellm_call_id="test-call-once",
|
||||
function_id="test-function-once",
|
||||
)
|
||||
test_logging_obj.handle_sync_success_callbacks_for_async_calls = MagicMock()
|
||||
|
||||
await litellm.acompletion(
|
||||
model="openai/codex-mini-latest",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
max_tokens=10,
|
||||
mock_response="hi",
|
||||
caching=False,
|
||||
litellm_logging_obj=test_logging_obj,
|
||||
)
|
||||
await asyncio.sleep(1.0)
|
||||
|
||||
assert (
|
||||
test_logging_obj.handle_sync_success_callbacks_for_async_calls.call_count == 1
|
||||
)
|
||||
|
||||
|
||||
def test_success_handler_does_not_emit_standard_payload_when_called_from_async(logging_obj):
|
||||
"""
|
||||
Regression test: sync success handler should not emit standard payload when
|
||||
invoked from async flow (async_success_handler already emits).
|
||||
"""
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
logging_obj.stream = True
|
||||
logging_obj.call_type = CallTypes.anthropic_messages.value
|
||||
logging_obj.model_call_details["litellm_params"] = {}
|
||||
logging_obj.litellm_params = {}
|
||||
|
||||
model_response = ModelResponse(
|
||||
id="resp-123",
|
||||
model="claude-sonnet-4-5",
|
||||
choices=[
|
||||
{
|
||||
"message": {"role": "assistant", "content": "hello"},
|
||||
"finish_reason": "stop",
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
usage={"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.litellm_core_utils.litellm_logging.emit_standard_logging_payload"
|
||||
) as mock_emit:
|
||||
logging_obj.success_handler(result=model_response, called_from_async=True)
|
||||
|
||||
mock_emit.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("call_type", ["completion", "responses"])
|
||||
def test_success_handler_runs_sync_callbacks_for_sync_requests(logging_obj, call_type):
|
||||
"""Ensure sync success callbacks execute when call type is sync (completion/responses)."""
|
||||
|
|
|
|||
|
|
@ -22,6 +22,7 @@ from litellm.litellm_core_utils.streaming_handler import (
|
|||
from litellm.types.utils import (
|
||||
CompletionTokensDetailsWrapper,
|
||||
Delta,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
PromptTokensDetailsWrapper,
|
||||
StandardLoggingPayload,
|
||||
|
|
@ -1355,6 +1356,156 @@ def test_usage_chunk_after_finish_reason_updates_hidden_params(logging_obj):
|
|||
f"Expected completion_tokens=135 from provider, got {hidden_usage.completion_tokens}"
|
||||
)
|
||||
|
||||
|
||||
@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()
|
||||
|
||||
bg_task = asyncio.create_task(background_tick())
|
||||
anext_task = asyncio.create_task(wrapper.__anext__())
|
||||
try:
|
||||
# If the event loop is blocked by a sync next(), this will time out.
|
||||
await asyncio.wait_for(tick_event.wait(), timeout=0.15)
|
||||
|
||||
out = await asyncio.wait_for(anext_task, timeout=2.0)
|
||||
assert isinstance(out, ModelResponseStream)
|
||||
finally:
|
||||
if not anext_task.done():
|
||||
anext_task.cancel()
|
||||
try:
|
||||
await anext_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
await bg_task
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_custom_stream_wrapper_anext_marks_sync_success_handler_as_async_origin():
|
||||
"""
|
||||
Regression test: async stream finalization should call success_handler with
|
||||
called_from_async=True to avoid duplicate standard payload emission.
|
||||
"""
|
||||
|
||||
class EmptyAsyncIterator:
|
||||
def __aiter__(self):
|
||||
return self
|
||||
|
||||
async def __anext__(self):
|
||||
raise StopAsyncIteration
|
||||
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.model_call_details = {"litellm_params": {}}
|
||||
logging_obj.async_success_handler = AsyncMock(return_value=None)
|
||||
|
||||
wrapper = CustomStreamWrapper(
|
||||
completion_stream=EmptyAsyncIterator(),
|
||||
model="vertex_ai/gemini-2.5-pro",
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
wrapper.sent_last_chunk = True
|
||||
wrapper.chunks = [
|
||||
ModelResponseStream(
|
||||
id="chunk-id",
|
||||
created=int(time.time()),
|
||||
model="vertex_ai/gemini-2.5-pro",
|
||||
object="chat.completion.chunk",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
delta=Delta(content="hello", role="assistant"),
|
||||
)
|
||||
],
|
||||
)
|
||||
]
|
||||
|
||||
complete_streaming_response = ModelResponse(
|
||||
id="resp-123",
|
||||
model="vertex_ai/gemini-2.5-pro",
|
||||
choices=[
|
||||
{
|
||||
"message": {"role": "assistant", "content": "hello"},
|
||||
"finish_reason": "stop",
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
usage={"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.litellm_core_utils.streaming_handler.litellm.stream_chunk_builder",
|
||||
return_value=complete_streaming_response,
|
||||
),
|
||||
patch(
|
||||
"litellm.litellm_core_utils.streaming_handler.executor.submit"
|
||||
) as mock_submit,
|
||||
):
|
||||
with pytest.raises(StopAsyncIteration):
|
||||
await wrapper.__anext__()
|
||||
|
||||
assert mock_submit.call_count == 1
|
||||
assert mock_submit.call_args.kwargs.get("called_from_async") is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_custom_stream_wrapper_aclose():
|
||||
"""Test that aclose() delegates to the underlying completion_stream's aclose()"""
|
||||
|
|
|
|||
119
tests/test_litellm/proxy/auth/test_info_routes.py
Normal file
119
tests/test_litellm/proxy/auth/test_info_routes.py
Normal file
|
|
@ -0,0 +1,119 @@
|
|||
import pytest
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from fastapi import HTTPException, Request
|
||||
|
||||
from litellm.proxy._types import LiteLLM_UserTable, LiteLLMRoutes, LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
|
||||
|
||||
def test_info_route_identification():
|
||||
"""Test that info routes are correctly identified"""
|
||||
for route in LiteLLMRoutes.info_routes.value:
|
||||
assert RouteChecks.is_info_route(route) is True
|
||||
|
||||
# Non-info routes should return False
|
||||
assert RouteChecks.is_info_route("/chat/completions") is False
|
||||
assert RouteChecks.is_info_route("/key/generate") is False
|
||||
|
||||
|
||||
def test_key_info_route_access():
|
||||
"""Test access control for /key/info route"""
|
||||
# This route handles its own access control, so it should pass for any user
|
||||
user_obj = LiteLLM_UserTable(
|
||||
user_id="test_user",
|
||||
user_email="test@example.com",
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
)
|
||||
valid_token = UserAPIKeyAuth(user_id="test_user")
|
||||
request = MagicMock(spec=Request)
|
||||
request.query_params = {}
|
||||
|
||||
# Should not raise exception as /key/info handles its own logic
|
||||
RouteChecks.non_proxy_admin_allowed_routes_check(
|
||||
user_obj=user_obj,
|
||||
_user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
route="/key/info",
|
||||
request=request,
|
||||
valid_token=valid_token,
|
||||
request_data={},
|
||||
)
|
||||
|
||||
|
||||
def test_user_info_route_access():
|
||||
"""Test access control for /user/info route"""
|
||||
user_obj = LiteLLM_UserTable(
|
||||
user_id="test_user",
|
||||
user_email="test@example.com",
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
)
|
||||
valid_token = UserAPIKeyAuth(user_id="test_user")
|
||||
request = MagicMock(spec=Request)
|
||||
request.query_params = {"user_id": "test_user"}
|
||||
|
||||
# Should not raise exception when user_id matches token's user_id
|
||||
RouteChecks.non_proxy_admin_allowed_routes_check(
|
||||
user_obj=user_obj,
|
||||
_user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
route="/user/info",
|
||||
request=request,
|
||||
valid_token=valid_token,
|
||||
request_data={},
|
||||
)
|
||||
|
||||
# Should raise exception when user_id does not match
|
||||
request.query_params = {"user_id": "different_user"}
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
RouteChecks.non_proxy_admin_allowed_routes_check(
|
||||
user_obj=user_obj,
|
||||
_user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
route="/user/info",
|
||||
request=request,
|
||||
valid_token=valid_token,
|
||||
request_data={},
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
|
||||
def test_model_info_route_access():
|
||||
"""Test access control for /model/info route"""
|
||||
user_obj = LiteLLM_UserTable(
|
||||
user_id="test_user",
|
||||
user_email="test@example.com",
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
)
|
||||
valid_token = UserAPIKeyAuth(user_id="test_user")
|
||||
request = MagicMock(spec=Request)
|
||||
request.query_params = {}
|
||||
|
||||
# Should not raise exception as /model/info is accessible to show user's models
|
||||
RouteChecks.non_proxy_admin_allowed_routes_check(
|
||||
user_obj=user_obj,
|
||||
_user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
route="/model/info",
|
||||
request=request,
|
||||
valid_token=valid_token,
|
||||
request_data={},
|
||||
)
|
||||
|
||||
|
||||
def test_team_info_route_access():
|
||||
"""Test access control for /team/info route"""
|
||||
user_obj = LiteLLM_UserTable(
|
||||
user_id="test_user",
|
||||
user_email="test@example.com",
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
)
|
||||
valid_token = UserAPIKeyAuth(user_id="test_user")
|
||||
request = MagicMock(spec=Request)
|
||||
request.query_params = {}
|
||||
|
||||
# Should not raise exception as /team/info handles its own logic
|
||||
RouteChecks.non_proxy_admin_allowed_routes_check(
|
||||
user_obj=user_obj,
|
||||
_user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
route="/team/info",
|
||||
request=request,
|
||||
valid_token=valid_token,
|
||||
request_data={},
|
||||
)
|
||||
|
|
@ -6460,6 +6460,8 @@ class TestValidateKeyAliasFormat:
|
|||
_validate_key_alias_format("valid/alias")
|
||||
_validate_key_alias_format("a" * 255)
|
||||
_validate_key_alias_format("my-key-123")
|
||||
_validate_key_alias_format("user/user@example.com")
|
||||
_validate_key_alias_format("team/user@example.com")
|
||||
|
||||
def test_validate_key_alias_format_invalid(self):
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import _validate_key_alias_format
|
||||
|
|
@ -6472,7 +6474,7 @@ class TestValidateKeyAliasFormat:
|
|||
"!", # special char
|
||||
"-start", # non-alphanumeric start
|
||||
"end-", # non-alphanumeric end
|
||||
"invalid@char", # invalid char
|
||||
"invalid#char", # invalid char
|
||||
"a" * 256, # too long
|
||||
" leading",
|
||||
"trailing ",
|
||||
|
|
|
|||
|
|
@ -35,7 +35,7 @@ def test_proxy_only_error_true_for_llm_route():
|
|||
)
|
||||
|
||||
|
||||
def test_proxy_only_error_false_for_non_llm_route():
|
||||
def test_proxy_only_error_true_for_info_route():
|
||||
proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())
|
||||
assert (
|
||||
proxy_logging_obj._is_proxy_only_llm_api_error(
|
||||
|
|
@ -43,6 +43,18 @@ def test_proxy_only_error_false_for_non_llm_route():
|
|||
error_type=ProxyErrorTypes.auth_error,
|
||||
route="/key/info",
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
|
||||
def test_proxy_only_error_false_for_non_llm_non_info_route():
|
||||
proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())
|
||||
assert (
|
||||
proxy_logging_obj._is_proxy_only_llm_api_error(
|
||||
original_exception=Exception(),
|
||||
error_type=ProxyErrorTypes.auth_error,
|
||||
route="/key/generate",
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -165,6 +165,40 @@ class TestManagedWebSocketHandlerIntegration:
|
|||
assert handler.timeout == 30.0
|
||||
assert handler.custom_llm_provider == "test_provider"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_websocket_log_messages_marks_sync_success_handler_as_async_origin(
|
||||
self,
|
||||
):
|
||||
"""WebSocket logging should suppress duplicate standard payload emission."""
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from litellm.responses.streaming_iterator import ResponsesWebSocketStreaming
|
||||
|
||||
mock_logging_obj = MagicMock()
|
||||
mock_logging_obj.async_success_handler = AsyncMock()
|
||||
mock_logging_obj.success_handler = MagicMock()
|
||||
mock_logging_obj.model_call_details = {}
|
||||
|
||||
streaming = ResponsesWebSocketStreaming(
|
||||
websocket=MagicMock(),
|
||||
backend_ws=MagicMock(),
|
||||
logging_obj=mock_logging_obj,
|
||||
)
|
||||
streaming.messages = [{"type": "response.completed"}]
|
||||
streaming.input_messages = [{"role": "user", "content": "hello"}]
|
||||
|
||||
with patch(
|
||||
"litellm.responses.streaming_iterator._ws_executor.submit"
|
||||
) as mock_submit:
|
||||
await streaming._log_messages()
|
||||
|
||||
assert mock_logging_obj.model_call_details["messages"] == streaming.input_messages
|
||||
assert mock_submit.call_args.args == (
|
||||
mock_logging_obj.success_handler,
|
||||
streaming.messages,
|
||||
)
|
||||
assert mock_submit.call_args.kwargs.get("called_from_async") is True
|
||||
|
||||
|
||||
class TestChunkTransformation:
|
||||
"""Test chunk serialization and transformation for WebSocket streaming"""
|
||||
|
|
|
|||
|
|
@ -0,0 +1,110 @@
|
|||
import React from "react";
|
||||
import { screen, waitFor } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { renderWithProviders } from "../../../tests/test-utils";
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
import * as networking from "../networking";
|
||||
import AddAttachmentForm from "./add_attachment_form";
|
||||
import { Policy } from "./types";
|
||||
|
||||
vi.mock("../networking");
|
||||
|
||||
vi.mock("./impact_preview_alert", () => ({
|
||||
default: ({ impactResult }: { impactResult: any }) =>
|
||||
React.createElement("div", { "data-testid": "impact-preview" }, `${impactResult.affected_keys_count} keys`),
|
||||
}));
|
||||
|
||||
const makePolicy = (overrides: Partial<Policy> = {}): Policy => ({
|
||||
policy_id: "policy-id-1",
|
||||
policy_name: "test-policy",
|
||||
inherit: null,
|
||||
description: null,
|
||||
guardrails_add: [],
|
||||
guardrails_remove: [],
|
||||
condition: null,
|
||||
...overrides,
|
||||
});
|
||||
|
||||
const defaultProps = {
|
||||
visible: true,
|
||||
onClose: vi.fn(),
|
||||
onSuccess: vi.fn(),
|
||||
accessToken: "test-token",
|
||||
policies: [makePolicy({ policy_name: "policy-alpha" }), makePolicy({ policy_name: "policy-beta", policy_id: "id-2" })],
|
||||
createAttachment: vi.fn(),
|
||||
};
|
||||
|
||||
describe("AddAttachmentForm", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
vi.mocked(networking.teamListCall).mockResolvedValue([]);
|
||||
vi.mocked(networking.keyListCall).mockResolvedValue({ keys: [] });
|
||||
vi.mocked(networking.modelAvailableCall).mockResolvedValue({ data: [] });
|
||||
});
|
||||
|
||||
it("should render the modal title when visible", async () => {
|
||||
renderWithProviders(<AddAttachmentForm {...defaultProps} />);
|
||||
expect(await screen.findByText("Create Policy Attachment")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should not render modal content when visible is false", () => {
|
||||
renderWithProviders(<AddAttachmentForm {...defaultProps} visible={false} />);
|
||||
expect(screen.queryByText("Create Policy Attachment")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should fetch teams, keys, and models on mount when visible and accessToken are provided", async () => {
|
||||
renderWithProviders(<AddAttachmentForm {...defaultProps} />);
|
||||
await waitFor(() => {
|
||||
expect(networking.teamListCall).toHaveBeenCalled();
|
||||
expect(networking.keyListCall).toHaveBeenCalled();
|
||||
expect(networking.modelAvailableCall).toHaveBeenCalled();
|
||||
});
|
||||
});
|
||||
|
||||
it("should not fetch teams, keys, or models when accessToken is null", () => {
|
||||
renderWithProviders(<AddAttachmentForm {...defaultProps} accessToken={null} />);
|
||||
expect(networking.teamListCall).not.toHaveBeenCalled();
|
||||
expect(networking.keyListCall).not.toHaveBeenCalled();
|
||||
expect(networking.modelAvailableCall).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("should call onClose when the Cancel button is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<AddAttachmentForm {...defaultProps} />);
|
||||
await user.click(await screen.findByRole("button", { name: /cancel/i }));
|
||||
expect(defaultProps.onClose).toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("should not show scope-specific fields when scope is global (default)", async () => {
|
||||
renderWithProviders(<AddAttachmentForm {...defaultProps} />);
|
||||
await screen.findByText("Create Policy Attachment");
|
||||
expect(screen.queryByText("Teams")).not.toBeInTheDocument();
|
||||
expect(screen.queryByText("Keys")).not.toBeInTheDocument();
|
||||
expect(screen.queryByText("Models")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show Teams, Keys, Models, and Tags fields when scope is switched to specific", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<AddAttachmentForm {...defaultProps} />);
|
||||
await screen.findByText("Create Policy Attachment");
|
||||
await user.click(screen.getByRole("radio", { name: /specific/i }));
|
||||
expect(screen.getByText("Teams")).toBeInTheDocument();
|
||||
expect(screen.getByText("Keys")).toBeInTheDocument();
|
||||
expect(screen.getByText("Models")).toBeInTheDocument();
|
||||
expect(screen.getByText("Tags")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show the 'Estimate Impact' button only when scope is specific", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<AddAttachmentForm {...defaultProps} />);
|
||||
await screen.findByText("Create Policy Attachment");
|
||||
expect(screen.queryByRole("button", { name: /estimate impact/i })).not.toBeInTheDocument();
|
||||
await user.click(screen.getByRole("radio", { name: /specific/i }));
|
||||
expect(screen.getByRole("button", { name: /estimate impact/i })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should render a 'Create Attachment' submit button", async () => {
|
||||
renderWithProviders(<AddAttachmentForm {...defaultProps} />);
|
||||
expect(await screen.findByRole("button", { name: /create attachment/i })).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
|
@ -0,0 +1,127 @@
|
|||
import React from "react";
|
||||
import { screen } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { renderWithProviders } from "../../../tests/test-utils";
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
import GuardrailSelectionModal from "./guardrail_selection_modal";
|
||||
|
||||
const makeGuardrailDef = (name: string, description = "A guardrail description") => ({
|
||||
guardrail_name: name,
|
||||
guardrail_info: { description },
|
||||
litellm_params: { guardrail: "presidio", mode: "pre_call" },
|
||||
});
|
||||
|
||||
const makeTemplate = (guardrailDefs: any[] = [], overrides: any = {}) => ({
|
||||
title: "Test Template",
|
||||
guardrailDefinitions: guardrailDefs,
|
||||
...overrides,
|
||||
});
|
||||
|
||||
const defaultProps = {
|
||||
visible: true,
|
||||
template: makeTemplate([makeGuardrailDef("guardrail-new-1"), makeGuardrailDef("guardrail-new-2")]),
|
||||
existingGuardrails: new Set<string>(),
|
||||
onConfirm: vi.fn(),
|
||||
onCancel: vi.fn(),
|
||||
};
|
||||
|
||||
describe("GuardrailSelectionModal", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
});
|
||||
|
||||
it("should render guardrail names from the template", async () => {
|
||||
renderWithProviders(<GuardrailSelectionModal {...defaultProps} />);
|
||||
expect(await screen.findByText("guardrail-new-1")).toBeInTheDocument();
|
||||
expect(screen.getByText("guardrail-new-2")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should pre-select only new guardrails when the modal opens", async () => {
|
||||
renderWithProviders(<GuardrailSelectionModal {...defaultProps} />);
|
||||
await screen.findByText("guardrail-new-1");
|
||||
const checkboxes = screen.getAllByRole("checkbox");
|
||||
checkboxes.forEach((cb) => expect(cb).toBeChecked());
|
||||
});
|
||||
|
||||
it("should not show a checkbox for guardrails that already exist", async () => {
|
||||
const props = {
|
||||
...defaultProps,
|
||||
template: makeTemplate([makeGuardrailDef("existing-g"), makeGuardrailDef("new-g")]),
|
||||
existingGuardrails: new Set(["existing-g"]),
|
||||
};
|
||||
renderWithProviders(<GuardrailSelectionModal {...props} />);
|
||||
await screen.findByText("existing-g");
|
||||
expect(screen.getAllByRole("checkbox")).toHaveLength(1);
|
||||
});
|
||||
|
||||
it("should show an 'Already exists' tag for guardrails that exist in the system", async () => {
|
||||
const props = {
|
||||
...defaultProps,
|
||||
template: makeTemplate([makeGuardrailDef("existing-g")]),
|
||||
existingGuardrails: new Set(["existing-g"]),
|
||||
};
|
||||
renderWithProviders(<GuardrailSelectionModal {...props} />);
|
||||
expect(await screen.findByText("Already exists")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show 'Create N Guardrails & Use Template' on the confirm button when N guardrails are selected", async () => {
|
||||
renderWithProviders(<GuardrailSelectionModal {...defaultProps} />);
|
||||
expect(await screen.findByRole("button", { name: /create 2 guardrails & use template/i })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show 'Use Template' on the confirm button when no new guardrails are selected", async () => {
|
||||
const props = {
|
||||
...defaultProps,
|
||||
template: makeTemplate([makeGuardrailDef("existing-g")]),
|
||||
existingGuardrails: new Set(["existing-g"]),
|
||||
};
|
||||
renderWithProviders(<GuardrailSelectionModal {...props} />);
|
||||
expect(await screen.findByRole("button", { name: /^use template$/i })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should call onConfirm with the definitions of selected guardrails when confirmed", async () => {
|
||||
const user = userEvent.setup();
|
||||
const def = makeGuardrailDef("my-guardrail");
|
||||
const props = { ...defaultProps, template: makeTemplate([def]) };
|
||||
renderWithProviders(<GuardrailSelectionModal {...props} />);
|
||||
await user.click(await screen.findByRole("button", { name: /create 1 guardrail/i }));
|
||||
expect(defaultProps.onConfirm).toHaveBeenCalledWith([def]);
|
||||
});
|
||||
|
||||
it("should deselect all guardrails when 'Deselect All' is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<GuardrailSelectionModal {...defaultProps} />);
|
||||
await screen.findByText("guardrail-new-1");
|
||||
await user.click(screen.getByRole("button", { name: /deselect all/i }));
|
||||
screen.getAllByRole("checkbox").forEach((cb) => expect(cb).not.toBeChecked());
|
||||
});
|
||||
|
||||
it("should re-select all new guardrails when 'Select All New' is clicked after deselecting", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<GuardrailSelectionModal {...defaultProps} />);
|
||||
await screen.findByText("guardrail-new-1");
|
||||
await user.click(screen.getByRole("button", { name: /deselect all/i }));
|
||||
await user.click(screen.getByRole("button", { name: /select all new/i }));
|
||||
screen.getAllByRole("checkbox").forEach((cb) => expect(cb).toBeChecked());
|
||||
});
|
||||
|
||||
it("should show 'No guardrails defined' when the template has no guardrail definitions", async () => {
|
||||
const props = { ...defaultProps, template: makeTemplate([]) };
|
||||
renderWithProviders(<GuardrailSelectionModal {...props} />);
|
||||
expect(await screen.findByText(/no guardrails defined for this template/i)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show a progress badge when progressInfo is provided", async () => {
|
||||
const props = { ...defaultProps, progressInfo: { current: 2, total: 5 } };
|
||||
renderWithProviders(<GuardrailSelectionModal {...props} />);
|
||||
expect(await screen.findByText(/template 2 of 5/i)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should call onCancel when the Cancel button is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<GuardrailSelectionModal {...defaultProps} />);
|
||||
await screen.findByText("guardrail-new-1");
|
||||
await user.click(screen.getByRole("button", { name: /^cancel$/i }));
|
||||
expect(defaultProps.onCancel).toHaveBeenCalled();
|
||||
});
|
||||
});
|
||||
|
|
@ -0,0 +1,184 @@
|
|||
import React from "react";
|
||||
import { screen, waitFor } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { renderWithProviders } from "../../../tests/test-utils";
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
import * as networking from "../networking";
|
||||
import ImpactPopover from "./impact_popover";
|
||||
import { PolicyAttachment } from "./types";
|
||||
|
||||
vi.mock("../networking");
|
||||
|
||||
vi.mock("@heroicons/react/outline", () => ({
|
||||
EyeIcon: function EyeIcon() { return null; },
|
||||
}));
|
||||
|
||||
vi.mock("@tremor/react", async (importOriginal) => {
|
||||
const actual = await importOriginal<typeof import("@tremor/react")>();
|
||||
return {
|
||||
...actual,
|
||||
Icon: ({ icon: IconComp, onClick, className }: any) =>
|
||||
React.createElement("button", { type: "button", onClick, className }, IconComp?.displayName ?? IconComp?.name ?? "icon"),
|
||||
};
|
||||
});
|
||||
|
||||
// Expose the Popover's onOpenChange so tests can trigger it programmatically.
|
||||
vi.mock("antd", async (importOriginal) => {
|
||||
const actual = await importOriginal<any>();
|
||||
return {
|
||||
...actual,
|
||||
Popover: ({ children, onOpenChange, content }: any) =>
|
||||
React.createElement(
|
||||
"div",
|
||||
null,
|
||||
React.createElement("div", { "data-testid": "popover-content" }, content),
|
||||
React.createElement(
|
||||
"div",
|
||||
{
|
||||
role: "button",
|
||||
"aria-label": "open-popover",
|
||||
onClick: () => onOpenChange?.(true),
|
||||
},
|
||||
children
|
||||
)
|
||||
),
|
||||
Tooltip: ({ children }: any) => React.createElement(React.Fragment, null, children),
|
||||
Spin: () => React.createElement("span", null, "Loading..."),
|
||||
Tag: ({ children }: any) => React.createElement("span", null, children),
|
||||
};
|
||||
});
|
||||
|
||||
const makeAttachment = (overrides: Partial<PolicyAttachment> = {}): PolicyAttachment => ({
|
||||
attachment_id: "att-001",
|
||||
policy_name: "my-policy",
|
||||
scope: null,
|
||||
teams: [],
|
||||
keys: [],
|
||||
models: [],
|
||||
tags: [],
|
||||
...overrides,
|
||||
});
|
||||
|
||||
describe("ImpactPopover", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
});
|
||||
|
||||
it("should render", () => {
|
||||
renderWithProviders(<ImpactPopover attachment={makeAttachment()} accessToken="tok" />);
|
||||
expect(screen.getByRole("button", { name: /open-popover/i })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show 'Click to load' as the initial popover content", () => {
|
||||
renderWithProviders(<ImpactPopover attachment={makeAttachment()} accessToken="tok" />);
|
||||
expect(screen.getByText(/click to load/i)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should call estimateAttachmentImpactCall when the popover is opened", async () => {
|
||||
const user = userEvent.setup();
|
||||
vi.mocked(networking.estimateAttachmentImpactCall).mockResolvedValue({
|
||||
affected_keys_count: 0,
|
||||
affected_teams_count: 0,
|
||||
sample_keys: [],
|
||||
sample_teams: [],
|
||||
});
|
||||
const attachment = makeAttachment({ policy_name: "rate-limit", teams: ["team-a"] });
|
||||
renderWithProviders(<ImpactPopover attachment={attachment} accessToken="my-token" />);
|
||||
await user.click(screen.getByRole("button", { name: /open-popover/i }));
|
||||
await waitFor(() => {
|
||||
expect(networking.estimateAttachmentImpactCall).toHaveBeenCalledWith("my-token", {
|
||||
policy_name: "rate-limit",
|
||||
scope: null,
|
||||
teams: ["team-a"],
|
||||
keys: [],
|
||||
models: [],
|
||||
tags: [],
|
||||
});
|
||||
});
|
||||
});
|
||||
|
||||
it("should not call the API when accessToken is null", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<ImpactPopover attachment={makeAttachment()} accessToken={null} />);
|
||||
await user.click(screen.getByRole("button", { name: /open-popover/i }));
|
||||
expect(networking.estimateAttachmentImpactCall).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("should show a loading indicator while the impact is being fetched", async () => {
|
||||
const user = userEvent.setup();
|
||||
vi.mocked(networking.estimateAttachmentImpactCall).mockReturnValue(new Promise(() => {}));
|
||||
renderWithProviders(<ImpactPopover attachment={makeAttachment()} accessToken="tok" />);
|
||||
await user.click(screen.getByRole("button", { name: /open-popover/i }));
|
||||
// Multiple "Loading..." nodes exist (Spin + adjacent text) — assert at least one is present
|
||||
expect(screen.queryAllByText(/loading/i).length).toBeGreaterThan(0);
|
||||
});
|
||||
|
||||
it("should show a global scope warning when affected_keys_count is -1", async () => {
|
||||
const user = userEvent.setup();
|
||||
vi.mocked(networking.estimateAttachmentImpactCall).mockResolvedValue({
|
||||
affected_keys_count: -1,
|
||||
affected_teams_count: -1,
|
||||
sample_keys: [],
|
||||
sample_teams: [],
|
||||
});
|
||||
renderWithProviders(<ImpactPopover attachment={makeAttachment()} accessToken="tok" />);
|
||||
await user.click(screen.getByRole("button", { name: /open-popover/i }));
|
||||
expect(await screen.findByText(/global scope.*affects all keys and teams/i)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show key and team counts when impact data is loaded for a specific scope", async () => {
|
||||
const user = userEvent.setup();
|
||||
vi.mocked(networking.estimateAttachmentImpactCall).mockResolvedValue({
|
||||
affected_keys_count: 5,
|
||||
affected_teams_count: 2,
|
||||
sample_keys: ["sk-abc"],
|
||||
sample_teams: ["team-x"],
|
||||
});
|
||||
renderWithProviders(<ImpactPopover attachment={makeAttachment()} accessToken="tok" />);
|
||||
await user.click(screen.getByRole("button", { name: /open-popover/i }));
|
||||
expect(await screen.findByText(/5/)).toBeInTheDocument();
|
||||
expect(screen.getByText(/2/)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should render sample key tags when returned from the API", async () => {
|
||||
const user = userEvent.setup();
|
||||
vi.mocked(networking.estimateAttachmentImpactCall).mockResolvedValue({
|
||||
affected_keys_count: 2,
|
||||
affected_teams_count: 0,
|
||||
sample_keys: ["sk-key-one", "sk-key-two"],
|
||||
sample_teams: [],
|
||||
});
|
||||
renderWithProviders(<ImpactPopover attachment={makeAttachment()} accessToken="tok" />);
|
||||
await user.click(screen.getByRole("button", { name: /open-popover/i }));
|
||||
expect(await screen.findByText("sk-key-one")).toBeInTheDocument();
|
||||
expect(screen.getByText("sk-key-two")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show 'No keys or teams currently affected' when both counts are 0", async () => {
|
||||
const user = userEvent.setup();
|
||||
vi.mocked(networking.estimateAttachmentImpactCall).mockResolvedValue({
|
||||
affected_keys_count: 0,
|
||||
affected_teams_count: 0,
|
||||
sample_keys: [],
|
||||
sample_teams: [],
|
||||
});
|
||||
renderWithProviders(<ImpactPopover attachment={makeAttachment()} accessToken="tok" />);
|
||||
await user.click(screen.getByRole("button", { name: /open-popover/i }));
|
||||
expect(await screen.findByText(/no keys or teams currently affected/i)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should not call the API a second time when the popover is already loaded", async () => {
|
||||
const user = userEvent.setup();
|
||||
vi.mocked(networking.estimateAttachmentImpactCall).mockResolvedValue({
|
||||
affected_keys_count: 1,
|
||||
affected_teams_count: 0,
|
||||
sample_keys: ["sk-abc"],
|
||||
sample_teams: [],
|
||||
});
|
||||
renderWithProviders(<ImpactPopover attachment={makeAttachment()} accessToken="tok" />);
|
||||
await user.click(screen.getByRole("button", { name: /open-popover/i }));
|
||||
await screen.findByText("sk-abc");
|
||||
await user.click(screen.getByRole("button", { name: /open-popover/i }));
|
||||
expect(networking.estimateAttachmentImpactCall).toHaveBeenCalledTimes(1);
|
||||
});
|
||||
});
|
||||
|
|
@ -0,0 +1,150 @@
|
|||
import React from "react";
|
||||
import { screen } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { renderWithProviders } from "../../../tests/test-utils";
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
import PolicyTable from "./policy_table";
|
||||
import { Policy } from "./types";
|
||||
|
||||
vi.mock("@heroicons/react/outline", () => ({
|
||||
TrashIcon: function TrashIcon() { return null; },
|
||||
PencilIcon: function PencilIcon() { return null; },
|
||||
SwitchVerticalIcon: function SwitchVerticalIcon() { return null; },
|
||||
ChevronUpIcon: function ChevronUpIcon() { return null; },
|
||||
ChevronDownIcon: function ChevronDownIcon() { return null; },
|
||||
}));
|
||||
|
||||
vi.mock("@tremor/react", async (importOriginal) => {
|
||||
const actual = await importOriginal<typeof import("@tremor/react")>();
|
||||
return {
|
||||
...actual,
|
||||
Button: React.forwardRef<HTMLButtonElement, any>(({ children, ...props }, ref) =>
|
||||
React.createElement("button", { ...props, ref }, children)
|
||||
),
|
||||
Icon: ({ icon: IconComp, onClick, className }: any) =>
|
||||
React.createElement("button", { type: "button", onClick, className }, IconComp?.displayName ?? IconComp?.name ?? "icon"),
|
||||
Tooltip: ({ children }: { children?: React.ReactNode }) =>
|
||||
React.createElement(React.Fragment, null, children),
|
||||
Badge: ({ children }: { children?: React.ReactNode }) =>
|
||||
React.createElement("span", null, children),
|
||||
};
|
||||
});
|
||||
|
||||
const makePolicy = (overrides: Partial<Policy> = {}): Policy => ({
|
||||
policy_id: "policy-id-1",
|
||||
policy_name: "test-policy",
|
||||
inherit: null,
|
||||
description: null,
|
||||
guardrails_add: [],
|
||||
guardrails_remove: [],
|
||||
condition: null,
|
||||
...overrides,
|
||||
});
|
||||
|
||||
const defaultProps = {
|
||||
policies: [],
|
||||
isLoading: false,
|
||||
onDeleteClick: vi.fn(),
|
||||
onEditClick: vi.fn(),
|
||||
onViewClick: vi.fn(),
|
||||
isAdmin: true,
|
||||
};
|
||||
|
||||
describe("PolicyTable", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
});
|
||||
|
||||
it("should render column headers", () => {
|
||||
renderWithProviders(<PolicyTable {...defaultProps} />);
|
||||
expect(screen.getByText("Name")).toBeInTheDocument();
|
||||
expect(screen.getByText("Description")).toBeInTheDocument();
|
||||
expect(screen.getByText("Actions")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show a loading message when isLoading is true", () => {
|
||||
renderWithProviders(<PolicyTable {...defaultProps} isLoading />);
|
||||
expect(screen.getByText(/loading/i)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show 'No policies found' when there are no policies", () => {
|
||||
renderWithProviders(<PolicyTable {...defaultProps} />);
|
||||
expect(screen.getByText(/no policies found/i)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should render a button with the policy name for each grouped policy", () => {
|
||||
const policies = [
|
||||
makePolicy({ policy_name: "alpha-policy", policy_id: "id-1" }),
|
||||
makePolicy({ policy_name: "beta-policy", policy_id: "id-2" }),
|
||||
];
|
||||
renderWithProviders(<PolicyTable {...defaultProps} policies={policies} />);
|
||||
expect(screen.getByRole("button", { name: "alpha-policy" })).toBeInTheDocument();
|
||||
expect(screen.getByRole("button", { name: "beta-policy" })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should call onViewClick with the policy_id when the policy name button is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
const policy = makePolicy({ policy_name: "my-policy", policy_id: "view-id-1" });
|
||||
renderWithProviders(<PolicyTable {...defaultProps} policies={[policy]} />);
|
||||
await user.click(screen.getByRole("button", { name: "my-policy" }));
|
||||
expect(defaultProps.onViewClick).toHaveBeenCalledWith("view-id-1");
|
||||
});
|
||||
|
||||
it("should call onDeleteClick with policy_id and policy_name when the delete icon is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
const policy = makePolicy({ policy_name: "del-policy", policy_id: "del-id-1" });
|
||||
renderWithProviders(<PolicyTable {...defaultProps} policies={[policy]} />);
|
||||
await user.click(screen.getByRole("button", { name: /TrashIcon/i }));
|
||||
expect(defaultProps.onDeleteClick).toHaveBeenCalledWith("del-id-1", "del-policy");
|
||||
});
|
||||
|
||||
it("should call onEditClick with the policy when the edit icon is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
const policy = makePolicy({ policy_name: "edit-policy", policy_id: "edit-id-1" });
|
||||
renderWithProviders(<PolicyTable {...defaultProps} policies={[policy]} />);
|
||||
await user.click(screen.getByRole("button", { name: /PencilIcon/i }));
|
||||
expect(defaultProps.onEditClick).toHaveBeenCalledWith(policy);
|
||||
});
|
||||
|
||||
it("should not show admin action icons for non-admins", () => {
|
||||
const policy = makePolicy();
|
||||
renderWithProviders(<PolicyTable {...defaultProps} policies={[policy]} isAdmin={false} />);
|
||||
expect(screen.queryByRole("button", { name: /TrashIcon/i })).not.toBeInTheDocument();
|
||||
expect(screen.queryByRole("button", { name: /PencilIcon/i })).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show a version badge when multiple versions of the same policy name exist", () => {
|
||||
const policies = [
|
||||
makePolicy({ policy_name: "versioned", policy_id: "v1", version_status: "published", version_number: 1 }),
|
||||
makePolicy({ policy_name: "versioned", policy_id: "v2", version_status: "production", version_number: 2 }),
|
||||
];
|
||||
renderWithProviders(<PolicyTable {...defaultProps} policies={policies} />);
|
||||
expect(screen.getByText(/2 version/i)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should group policies with the same name into a single row", () => {
|
||||
const policies = [
|
||||
makePolicy({ policy_name: "shared", policy_id: "s1", version_status: "published" }),
|
||||
makePolicy({ policy_name: "shared", policy_id: "s2", version_status: "production" }),
|
||||
];
|
||||
renderWithProviders(<PolicyTable {...defaultProps} policies={policies} />);
|
||||
expect(screen.getAllByRole("button", { name: "shared" })).toHaveLength(1);
|
||||
});
|
||||
|
||||
it("should show an overflow tag when more than 2 guardrails_add exist", () => {
|
||||
const policy = makePolicy({ guardrails_add: ["g1", "g2", "g3", "g4"] });
|
||||
renderWithProviders(<PolicyTable {...defaultProps} policies={[policy]} />);
|
||||
expect(screen.getByText("+2")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should prefer the production version as the primary policy when grouping", async () => {
|
||||
const user = userEvent.setup();
|
||||
const policies = [
|
||||
makePolicy({ policy_name: "grouped", policy_id: "published-id", version_status: "published" }),
|
||||
makePolicy({ policy_name: "grouped", policy_id: "prod-id", version_status: "production" }),
|
||||
];
|
||||
renderWithProviders(<PolicyTable {...defaultProps} policies={policies} />);
|
||||
await user.click(screen.getByRole("button", { name: "grouped" }));
|
||||
expect(defaultProps.onViewClick).toHaveBeenCalledWith("prod-id");
|
||||
});
|
||||
});
|
||||
|
|
@ -0,0 +1,144 @@
|
|||
import React from "react";
|
||||
import { screen, waitFor } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { renderWithProviders } from "../../../tests/test-utils";
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
import * as networking from "../networking";
|
||||
import PolicyTemplates from "./policy_templates";
|
||||
|
||||
vi.mock("../networking");
|
||||
|
||||
vi.mock("@heroicons/react/outline", () => ({
|
||||
ShieldCheckIcon: function ShieldCheckIcon() { return null; },
|
||||
ShieldExclamationIcon: function ShieldExclamationIcon() { return null; },
|
||||
BeakerIcon: function BeakerIcon() { return null; },
|
||||
CurrencyDollarIcon: function CurrencyDollarIcon() { return null; },
|
||||
CheckCircleIcon: function CheckCircleIcon() { return null; },
|
||||
}));
|
||||
|
||||
const makeTemplate = (overrides: any = {}) => ({
|
||||
id: "tpl-1",
|
||||
title: "Test Template",
|
||||
description: "A test template",
|
||||
icon: "ShieldCheckIcon",
|
||||
iconColor: "text-green-500",
|
||||
iconBg: "bg-green-50",
|
||||
guardrails: ["guardrail-a"],
|
||||
tags: [],
|
||||
complexity: "Low" as const,
|
||||
...overrides,
|
||||
});
|
||||
|
||||
const defaultProps = {
|
||||
onUseTemplate: vi.fn(),
|
||||
onOpenAiSuggestion: vi.fn(),
|
||||
accessToken: "test-token",
|
||||
};
|
||||
|
||||
describe("PolicyTemplates", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
});
|
||||
|
||||
it("should render the section header after loading", async () => {
|
||||
vi.mocked(networking.getPolicyTemplates).mockResolvedValue([]);
|
||||
renderWithProviders(<PolicyTemplates {...defaultProps} />);
|
||||
expect(await screen.findByText("Policy Templates")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should not show the template grid while fetching", () => {
|
||||
vi.mocked(networking.getPolicyTemplates).mockReturnValue(new Promise(() => {}));
|
||||
renderWithProviders(<PolicyTemplates {...defaultProps} />);
|
||||
expect(screen.queryByText("Policy Templates")).not.toBeInTheDocument();
|
||||
expect(screen.queryByRole("button", { name: /use ai to find templates/i })).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should render a card for each fetched template", async () => {
|
||||
const templates = [
|
||||
makeTemplate({ title: "Template Alpha" }),
|
||||
makeTemplate({ id: "tpl-2", title: "Template Beta" }),
|
||||
];
|
||||
vi.mocked(networking.getPolicyTemplates).mockResolvedValue(templates);
|
||||
renderWithProviders(<PolicyTemplates {...defaultProps} />);
|
||||
expect(await screen.findByText("Template Alpha")).toBeInTheDocument();
|
||||
expect(screen.getByText("Template Beta")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should call onTemplatesLoaded with the fetched templates after loading", async () => {
|
||||
const templates = [makeTemplate()];
|
||||
vi.mocked(networking.getPolicyTemplates).mockResolvedValue(templates);
|
||||
const onTemplatesLoaded = vi.fn();
|
||||
renderWithProviders(<PolicyTemplates {...defaultProps} onTemplatesLoaded={onTemplatesLoaded} />);
|
||||
await waitFor(() => {
|
||||
expect(onTemplatesLoaded).toHaveBeenCalledWith(templates);
|
||||
});
|
||||
});
|
||||
|
||||
it("should call onOpenAiSuggestion when the AI suggestion button is clicked", async () => {
|
||||
vi.mocked(networking.getPolicyTemplates).mockResolvedValue([]);
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<PolicyTemplates {...defaultProps} />);
|
||||
await screen.findByText("Policy Templates");
|
||||
await user.click(screen.getByRole("button", { name: /use ai to find templates/i }));
|
||||
expect(defaultProps.onOpenAiSuggestion).toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("should render tag filter checkboxes for unique tags across all templates", async () => {
|
||||
const templates = [
|
||||
makeTemplate({ tags: ["compliance"] }),
|
||||
makeTemplate({ id: "tpl-2", tags: ["compliance", "security"] }),
|
||||
];
|
||||
vi.mocked(networking.getPolicyTemplates).mockResolvedValue(templates);
|
||||
renderWithProviders(<PolicyTemplates {...defaultProps} />);
|
||||
expect(await screen.findByRole("checkbox", { name: /compliance/i })).toBeInTheDocument();
|
||||
expect(screen.getByRole("checkbox", { name: /security/i })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should filter to only matching templates when a tag is selected", async () => {
|
||||
const templates = [
|
||||
makeTemplate({ id: "tpl-1", title: "Compliance Template", tags: ["compliance"] }),
|
||||
makeTemplate({ id: "tpl-2", title: "Security Template", tags: ["security"] }),
|
||||
];
|
||||
vi.mocked(networking.getPolicyTemplates).mockResolvedValue(templates);
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<PolicyTemplates {...defaultProps} />);
|
||||
await screen.findByText("Compliance Template");
|
||||
await user.click(screen.getByRole("checkbox", { name: /compliance/i }));
|
||||
expect(screen.getByText("Compliance Template")).toBeInTheDocument();
|
||||
expect(screen.queryByText("Security Template")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show 'No templates match' when selected tags exclude all templates", async () => {
|
||||
const templates = [
|
||||
makeTemplate({ id: "tpl-1", title: "Alpha Template", tags: ["alpha"] }),
|
||||
makeTemplate({ id: "tpl-2", title: "Beta Template", tags: ["beta"] }),
|
||||
];
|
||||
vi.mocked(networking.getPolicyTemplates).mockResolvedValue(templates);
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<PolicyTemplates {...defaultProps} />);
|
||||
await screen.findByText("Alpha Template");
|
||||
await user.click(screen.getByRole("checkbox", { name: /alpha/i }));
|
||||
await user.click(screen.getByRole("checkbox", { name: /beta/i }));
|
||||
expect(screen.getByText(/no templates match the selected filters/i)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should restore all templates when 'Clear all' is clicked", async () => {
|
||||
const templates = [
|
||||
makeTemplate({ id: "tpl-1", title: "Alpha Template", tags: ["alpha"] }),
|
||||
makeTemplate({ id: "tpl-2", title: "Beta Template", tags: ["beta"] }),
|
||||
];
|
||||
vi.mocked(networking.getPolicyTemplates).mockResolvedValue(templates);
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<PolicyTemplates {...defaultProps} />);
|
||||
await screen.findByText("Alpha Template");
|
||||
await user.click(screen.getByRole("checkbox", { name: /alpha/i }));
|
||||
expect(screen.queryByText("Beta Template")).not.toBeInTheDocument();
|
||||
await user.click(screen.getByRole("button", { name: /clear all/i }));
|
||||
expect(screen.getByText("Beta Template")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should not fetch templates when accessToken is null", () => {
|
||||
renderWithProviders(<PolicyTemplates {...defaultProps} accessToken={null} />);
|
||||
expect(networking.getPolicyTemplates).not.toHaveBeenCalled();
|
||||
});
|
||||
});
|
||||
Loading…
Add table
Reference in a new issue