From 3250ccc802b30cd768e595347b7b2e15c9beef8b Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Fri, 9 Oct 2026 09:48:10 -0700 Subject: [PATCH] test: move offline logging, secret manager and provider tests from legacy dirs into tests/unit (#45554) Moves 58 legacy nodes (cyberark, logging callback manager, langfuse handler and helpers, SQS, app crypto, bedrock nova and invoke, cloudflare, litellm_proxy, nvidia nim, replicate, triton, convert_dict_to_response, get_model_info, vcr live-call probe) into the tests/unit files that mirror the modules they exercise. Calls that reached huggingface.co, langfuse or 0.0.0.0 are refused at the HTTP boundary with respx. Tests that need credentials, a license, a spawned server, sleeps, a subprocess, or assert nothing stay in their legacy files --- litellm-rust/crates/secrets/PARITY.md | 2 +- tests/litellm_utils_tests/test_cyberark.py | 404 ------------- .../test_logging_callback_manager.py | 41 -- .../test_bedrock_nova_embedding.py | 97 ---- tests/llm_translation/test_cloudflare.py | 99 ---- .../test_litellm_proxy_provider.py | 46 -- .../test_convert_dict_to_chat_completion.py | 93 --- tests/llm_translation/test_nvidia_nim.py | 72 --- tests/llm_translation/test_replicate.py | 103 ---- tests/llm_translation/test_triton.py | 100 ---- .../test_unit_test_bedrock_invoke.py | 34 -- .../test_vcr_classification.py | 148 ----- tests/local_testing/test_get_model_info.py | 34 -- .../test_langfuse_dynamic_credentials.py | 32 - .../test_langfuse_unit_tests.py | 547 ------------------ .../logging_callback_tests/test_sqs_logger.py | 291 ---------- .../test_langfuse_dynamic_credentials.py | 20 + .../langfuse/test_langfuse_handler.py | 163 ++++++ tests/unit/integrations/test_langfuse.py | 192 ++++++ tests/unit/integrations/test_sqs.py | 237 ++++++++ .../test_convert_dict_to_response.py | 75 +++ .../litellm_core_utils/test_app_crypto.py | 20 + .../test_litellm_logging.py | 42 ++ .../test_logging_callback_manager.py | 26 + .../test_model_param_helper.py | 27 + .../test_base_invoke_transformation.py | 19 + .../embed/test_amazon_nova_transformation.py | 17 + .../test_cloudflare_transformation.py | 64 ++ .../test_litellm_proxy_chat_transformation.py | 16 + tests/unit/llms/nvidia_nim/test_nvidia_nim.py | 40 ++ .../replicate/chat/test_transformation.py | 89 +++ tests/unit/llms/triton/test_triton.py | 113 ++++ .../test_cyberark_secret_manager.py | 281 +++++++++ tests/unit/test_utils_get_model_info.py | 31 + tests/unit/test_vcr_classification.py | 32 + 35 files changed, 1505 insertions(+), 2142 deletions(-) delete mode 100644 tests/litellm_utils_tests/test_cyberark.py delete mode 100644 tests/litellm_utils_tests/test_logging_callback_manager.py delete mode 100644 tests/llm_translation/test_bedrock_nova_embedding.py delete mode 100644 tests/llm_translation/test_cloudflare.py delete mode 100644 tests/llm_translation/test_litellm_proxy_provider.py delete mode 100644 tests/llm_translation/test_llm_response_utils/test_convert_dict_to_chat_completion.py delete mode 100644 tests/llm_translation/test_nvidia_nim.py delete mode 100644 tests/llm_translation/test_replicate.py delete mode 100644 tests/llm_translation/test_unit_test_bedrock_invoke.py delete mode 100644 tests/llm_translation/test_vcr_classification.py delete mode 100644 tests/local_testing/test_get_model_info.py delete mode 100644 tests/logging_callback_tests/test_langfuse_dynamic_credentials.py create mode 100644 tests/unit/integrations/langfuse/test_langfuse_handler.py create mode 100644 tests/unit/integrations/test_sqs.py create mode 100644 tests/unit/litellm_core_utils/test_app_crypto.py diff --git a/litellm-rust/crates/secrets/PARITY.md b/litellm-rust/crates/secrets/PARITY.md index 9439da93d65..ed9a8b29618 100644 --- a/litellm-rust/crates/secrets/PARITY.md +++ b/litellm-rust/crates/secrets/PARITY.md @@ -229,7 +229,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo | `test_hashicorp_secret_manager_rotate_secret_with_team_overrides` | [rotation_applies_timeout_to_each_request](../secrets-hashicorp/tests/secret_manager/writes.rs) | | `test_hashicorp_secret_manager_rotate_secret_value_mismatch` | [a_different_replacement_never_deletes_the_current_secret](../secrets-types/tests/rotation.rs) | -## [tests/litellm_utils_tests/test_cyberark.py](../../../tests/litellm_utils_tests/test_cyberark.py) +## [tests/unit/secret_managers/test_cyberark_secret_manager.py](../../../tests/unit/secret_managers/test_cyberark_secret_manager.py) (moved from tests/litellm_utils_tests/test_cyberark.py) | Python test | Rust coverage or boundary | | --- | --- | diff --git a/tests/litellm_utils_tests/test_cyberark.py b/tests/litellm_utils_tests/test_cyberark.py deleted file mode 100644 index 84b1fc5ebd2..00000000000 --- a/tests/litellm_utils_tests/test_cyberark.py +++ /dev/null @@ -1,404 +0,0 @@ -""" -Integration test for CyberArk Conjur Secret Manager. -""" - -import os -import pytest -import yaml -from dotenv import load_dotenv - -load_dotenv() - -from unittest.mock import AsyncMock, MagicMock, patch -from litellm._uuid import uuid - -# Set up environment variables for testing -os.environ["CYBERARK_API_KEY"] = "test-cyberark-api-key-909" -os.environ["CYBERARK_API_BASE"] = "http://0.0.0.0:8080" -os.environ["CYBERARK_ACCOUNT"] = "default" -os.environ["CYBERARK_USERNAME"] = "admin" - -from litellm.secret_managers.cyberark_secret_manager import CyberArkSecretManager - - -def create_mock_response(status_code: int, text: str = ""): - """ - Helper function to create a mock HTTP response. - """ - mock_response = MagicMock() - mock_response.status_code = status_code - mock_response.text = text - mock_response.raise_for_status = MagicMock() - - if status_code >= 400: - import httpx - - error = httpx.HTTPStatusError( - message=f"HTTP {status_code}", request=MagicMock(), response=mock_response - ) - mock_response.raise_for_status.side_effect = error - - return mock_response - - -@pytest.mark.asyncio -async def test_cyberark_write_secret_rejects_yaml_injection(): - """ - Regression test: async_write_secret must reject a secret_name that is not - safe to embed in the Conjur policy body, before any HTTP call is made. - """ - with patch("litellm.proxy.proxy_server.premium_user", True): - malicious_secret_name = "foo\n- !grant\n role: !!admin\n member: attacker" - - mock_sync_client = MagicMock() - mock_async_client = AsyncMock() - - with ( - patch( - "litellm.secret_managers.cyberark_secret_manager.get_httpx_client", - return_value=mock_sync_client, - ), - patch( - "litellm.secret_managers.cyberark_secret_manager.get_async_httpx_client", - return_value=mock_async_client, - ), - ): - cyberark_manager = CyberArkSecretManager() - - response = await cyberark_manager.async_write_secret( - secret_name=malicious_secret_name, - secret_value="sk-9876", - ) - - assert response["status"] == "error" - assert "Invalid secret_name" in response["message"] - # The malicious policy YAML must never reach the wire. - mock_sync_client.client.post.assert_not_called() - mock_async_client.post.assert_not_called() - - -@pytest.mark.parametrize( - "secret_name", - [ - "foo: bar", - "foo # bar", - "plain-alias", - "team/user@example.com", - ], -) -@pytest.mark.asyncio -async def test_cyberark_ensure_variable_exists_escapes_yaml_metacharacters(secret_name): - """ - Regression test: _ensure_variable_exists must escape secret_name (not just - denylist-check it) so the policy body always parses back to exactly one - '!variable' scalar node holding the untouched secret_name. - """ - with patch("litellm.proxy.proxy_server.premium_user", True): - captured = {} - - async def _capture_post(url, headers=None, content=None): - captured["content"] = content - return create_mock_response(status_code=201, text="") - - mock_sync_client = MagicMock() - mock_sync_client.client.post.return_value = create_mock_response(status_code=200, text="mock-token") - mock_async_client = MagicMock() - mock_async_client.client.post.side_effect = _capture_post - - with patch( - "litellm.secret_managers.cyberark_secret_manager.get_httpx_client", - return_value=mock_sync_client, - ): - cyberark_manager = CyberArkSecretManager() - await cyberark_manager._ensure_variable_exists(secret_name, mock_async_client) - - policy_yaml = captured["content"] - parsed = yaml.compose(policy_yaml) - assert len(parsed.value) == 1 - node = parsed.value[0] - assert node.tag == "!variable" - assert node.value == secret_name - - -@pytest.mark.asyncio -async def test_cyberark_write_and_read_secret(): - """ - Test writing a secret to CyberArk Conjur and reading it back using mocked HTTP requests. - """ - with patch("litellm.proxy.proxy_server.premium_user", True): - # Generate unique secret name and value - secret_name = f"test-secret-{uuid.uuid4()}" - secret_value = f"test-value-{uuid.uuid4()}" - - # Mock sync httpx client (for auth, ensure variable exists, sync read) - # The get_httpx_client returns an HTTPHandler with a .client property - mock_sync_client = MagicMock() - # Auth response - note: the actual client is accessed via .client property - mock_sync_client.client.post.return_value = create_mock_response(status_code=200, text="mock-token") - # Sync read response - mock_sync_client.client.get.return_value = create_mock_response(status_code=200, text=secret_value) - - # Mock async httpx client (for async write) - mock_async_client = AsyncMock() - mock_async_client.post.return_value = create_mock_response(status_code=201, text="") - - with ( - patch( - "litellm.secret_managers.cyberark_secret_manager.get_httpx_client", - return_value=mock_sync_client, - ), - patch( - "litellm.secret_managers.cyberark_secret_manager.get_async_httpx_client", - return_value=mock_async_client, - ), - ): - # Create CyberArk secret manager instance - cyberark_manager = CyberArkSecretManager() - - # Write the secret - write_response = await cyberark_manager.async_write_secret( - secret_name=secret_name, - secret_value=secret_value, - ) - print("write_response=", write_response) - - # Validate write was successful - assert write_response["status"] == "success" - - # Read the secret back - read_value = cyberark_manager.sync_read_secret(secret_name=secret_name) - print("READ VALUE=", read_value) - - # Validate the secret exists and has the correct value - assert read_value is not None - assert read_value == secret_value - - -@pytest.mark.asyncio -async def test_cyberark_rotate_secret(): - """ - Test key rotation in CyberArk Conjur using mocked HTTP requests. - - This test simulates what happens when a virtual key is rotated: - 1. Write initial secret with alias (like a proxy key) - 2. Rotate to new value (like sk-12359) - 3. Verify reading the secret returns the NEW value - """ - with patch("litellm.proxy.proxy_server.premium_user", True): - # Simulate initial virtual key creation - secret_alias = f"test-rotation-key-{uuid.uuid4()}" - initial_key_value = f"sk-initial-{uuid.uuid4()}" - rotated_key_value = f"sk-rotated-{uuid.uuid4()}" - - print(f"\n=== Testing Key Rotation ===") - print(f"Alias: {secret_alias}") - print(f"Initial value: {initial_key_value}") - print(f"Rotated value: {rotated_key_value}") - - # Store the current value to simulate actual storage behavior - current_value = {"value": initial_key_value} - - # Mock sync httpx client (for auth, ensure variable exists, sync reads) - # The get_httpx_client returns an HTTPHandler with a .client property - mock_sync_client = MagicMock() - # Auth response - note: the actual client is accessed via .client property - mock_sync_client.client.post.return_value = create_mock_response( - status_code=200, text="mock-token" - ) - - # Sync reads return the current value from our simulated storage - def get_mock_sync_read_response(*args, **kwargs): - return create_mock_response(status_code=200, text=current_value["value"]) - - mock_sync_client.client.get.side_effect = get_mock_sync_read_response - - # Mock async httpx client (for async writes and reads) - mock_async_client = AsyncMock() - - # Async writes update the current value - async def mock_async_post(*args, **kwargs): - content = kwargs.get("content", "") - if content: - current_value["value"] = content - return create_mock_response(status_code=201, text="") - - mock_async_client.post.side_effect = mock_async_post - - # Async reads also return the current value - async def get_mock_async_read_response(*args, **kwargs): - return create_mock_response(status_code=200, text=current_value["value"]) - - mock_async_client.get.side_effect = get_mock_async_read_response - - with ( - patch( - "litellm.secret_managers.cyberark_secret_manager.get_httpx_client", - return_value=mock_sync_client, - ), - patch( - "litellm.secret_managers.cyberark_secret_manager.get_async_httpx_client", - return_value=mock_async_client, - ), - ): - # Create CyberArk secret manager instance - cyberark_manager = CyberArkSecretManager() - - # Step 1: Write initial secret (simulates key creation) - write_response = await cyberark_manager.async_write_secret( - secret_name=secret_alias, - secret_value=initial_key_value, - ) - print(f"\n1. Initial write response: {write_response}") - assert write_response["status"] == "success" - - # Verify initial value was written - initial_read = cyberark_manager.sync_read_secret(secret_name=secret_alias) - print(f"2. Initial read value: {initial_read}") - assert initial_read == initial_key_value - - # Step 2: Rotate the secret (simulates key rotation) - # In key rotation, we keep the same secret_name but update the value - rotation_response = await cyberark_manager.async_rotate_secret( - current_secret_name=secret_alias, - new_secret_name=secret_alias, - new_secret_value=rotated_key_value, - ) - print(f"3. Rotation response: {rotation_response}") - assert rotation_response["status"] == "success" - - # Clear cache to force a fresh read - cyberark_manager.cache.flush_cache() - - # Step 3: Verify the secret now returns the NEW value - rotated_read = cyberark_manager.sync_read_secret(secret_name=secret_alias) - print(f"4. After rotation, read value: {rotated_read}") - - # This is the key assertion: after rotation, reading should return the NEW value - assert rotated_read is not None - assert rotated_read == rotated_key_value - assert rotated_read != initial_key_value - - print( - f"\n✅ Rotation successful: {initial_key_value} → {rotated_key_value}" - ) - - -@pytest.mark.asyncio -async def test_cyberark_rotate_secret_with_new_alias(): - """ - Test key rotation with a new alias using mocked HTTP requests. - - This simulates rotating a key and changing its alias at the same time: - 1. Write secret with alias-v1 - 2. Rotate to alias-v2 with new value - 3. Verify alias-v2 has the new value - 4. Verify alias-v1 still exists with old value (CyberArk doesn't delete) - """ - with patch("litellm.proxy.proxy_server.premium_user", True): - # Simulate key rotation with alias change - base_alias = f"test-alias-change-{uuid.uuid4()}" - old_alias = f"{base_alias}-v1" - new_alias = f"{base_alias}-v2" - old_value = f"sk-old-{uuid.uuid4()}" - new_value = f"sk-new-{uuid.uuid4()}" - - print(f"\n=== Testing Key Rotation with Alias Change ===") - print(f"Old alias: {old_alias} = {old_value}") - print(f"New alias: {new_alias} = {new_value}") - - # Store secrets in a dict to simulate actual storage - secrets_store = {} - - # Mock sync httpx client (for auth, ensure variable exists, sync reads) - # The get_httpx_client returns an HTTPHandler with a .client property - mock_sync_client = MagicMock() - # Auth response - note: the actual client is accessed via .client property - mock_sync_client.client.post.return_value = create_mock_response( - status_code=200, text="mock-token" - ) - - # Mock sync reads to return from our store - def get_mock_sync_read(*args, **kwargs): - url = args[0] if args else kwargs.get("url", "") - # Extract secret name from URL - for secret_name, secret_val in secrets_store.items(): - if secret_name in url: - return create_mock_response(status_code=200, text=secret_val) - return create_mock_response(status_code=404, text="Not found") - - mock_sync_client.client.get.side_effect = get_mock_sync_read - - # Mock async httpx client (for async writes and reads) - mock_async_client = AsyncMock() - - # Mock async write to update our store - async def mock_async_post(*args, **kwargs): - url = args[0] if args else kwargs.get("url", "") - content = kwargs.get("content", "") - - # Extract secret name from URL and store the value - if old_alias in url: - secrets_store[old_alias] = content - elif new_alias in url: - secrets_store[new_alias] = content - - return create_mock_response(status_code=201, text="") - - mock_async_client.post.side_effect = mock_async_post - - # Mock async reads to return from our store - async def get_mock_async_read(*args, **kwargs): - url = args[0] if args else kwargs.get("url", "") - # Extract secret name from URL - for secret_name, secret_val in secrets_store.items(): - if secret_name in url: - return create_mock_response(status_code=200, text=secret_val) - return create_mock_response(status_code=404, text="Not found") - - mock_async_client.get.side_effect = get_mock_async_read - - with ( - patch( - "litellm.secret_managers.cyberark_secret_manager.get_httpx_client", - return_value=mock_sync_client, - ), - patch( - "litellm.secret_managers.cyberark_secret_manager.get_async_httpx_client", - return_value=mock_async_client, - ), - ): - # Create CyberArk secret manager instance - cyberark_manager = CyberArkSecretManager() - - # Step 1: Create initial secret with old alias - write_response = await cyberark_manager.async_write_secret( - secret_name=old_alias, - secret_value=old_value, - ) - print(f"\n1. Initial write: {write_response}") - assert write_response["status"] == "success" - - # Step 2: Rotate to new alias with new value - rotation_response = await cyberark_manager.async_rotate_secret( - current_secret_name=old_alias, - new_secret_name=new_alias, - new_secret_value=new_value, - ) - print(f"2. Rotation response: {rotation_response}") - assert rotation_response["status"] == "success" - - # Clear cache to force fresh reads - cyberark_manager.cache.flush_cache() - - # Step 3: Verify new alias has new value - new_read = cyberark_manager.sync_read_secret(secret_name=new_alias) - print(f"3. Read new alias: {new_read}") - assert new_read == new_value - - # Step 4: Verify old alias still exists (CyberArk doesn't delete via API) - old_read = cyberark_manager.sync_read_secret(secret_name=old_alias) - print(f"4. Read old alias (should still exist): {old_read}") - assert old_read == old_value - - print(f"\n✅ Alias rotation successful: {old_alias} → {new_alias}") - print(f" Note: Old alias still exists in CyberArk (expected behavior)") diff --git a/tests/litellm_utils_tests/test_logging_callback_manager.py b/tests/litellm_utils_tests/test_logging_callback_manager.py deleted file mode 100644 index fdd1d8294f3..00000000000 --- a/tests/litellm_utils_tests/test_logging_callback_manager.py +++ /dev/null @@ -1,41 +0,0 @@ -import litellm -from litellm.litellm_core_utils.logging_callback_manager import LoggingCallbackManager -from litellm.integrations.langfuse.langfuse_prompt_management import ( - LangfusePromptManagement, -) -from litellm.integrations.opentelemetry import OpenTelemetry -def test_duplicate_langfuse_logger_test(): - manager = LoggingCallbackManager() - for _ in range(10): - langfuse_logger = LangfusePromptManagement() - manager.add_litellm_success_callback(langfuse_logger) - print("litellm.success_callback: ", litellm.success_callback) - assert len(litellm.success_callback) == 1 - - -def test_duplicate_multiple_loggers_test(): - manager = LoggingCallbackManager() - for _ in range(10): - langfuse_logger = LangfusePromptManagement() - otel_logger = OpenTelemetry() - manager.add_litellm_success_callback(langfuse_logger) - manager.add_litellm_success_callback(otel_logger) - print("litellm.success_callback: ", litellm.success_callback) - assert len(litellm.success_callback) == 2 - - # Check exactly one instance of each logger type - langfuse_count = sum( - 1 - for callback in litellm.success_callback - if isinstance(callback, LangfusePromptManagement) - ) - otel_count = sum( - 1 - for callback in litellm.success_callback - if isinstance(callback, OpenTelemetry) - ) - - assert ( - langfuse_count == 1 - ), "Should have exactly one LangfusePromptManagement instance" - assert otel_count == 1, "Should have exactly one OpenTelemetry instance" diff --git a/tests/llm_translation/test_bedrock_nova_embedding.py b/tests/llm_translation/test_bedrock_nova_embedding.py deleted file mode 100644 index b8c0a9ab8cc..00000000000 --- a/tests/llm_translation/test_bedrock_nova_embedding.py +++ /dev/null @@ -1,97 +0,0 @@ -""" -Test suite for Amazon Nova Multimodal Embeddings integration with LiteLLM. - -Tests cover: -- Synchronous text embeddings -- Synchronous image embeddings -- Synchronous video/audio embeddings -- Asynchronous embeddings with segmentation -- Different embedding purposes and dimensions -- Error handling -""" - - -import pytest - -from litellm.llms.bedrock.embed.amazon_nova_transformation import ( - AmazonNovaEmbeddingConfig, -) - - -class TestNovaTransformationRequest: - """Test request transformation for Nova embeddings.""" - - - - - - - def test_async_invoke_requires_output_s3_uri(self): - """Test that async invoke requires output_s3_uri.""" - config = AmazonNovaEmbeddingConfig() - - inference_params = { - "embedding_purpose": "GENERIC_INDEX", - } - - with pytest.raises(ValueError, match="output_s3_uri is required"): - config.transform_request( - input="Test text", - inference_params=inference_params, - async_invoke_route=True, - model_id="amazon.nova-2-multimodal-embeddings-v1:0", - output_s3_uri=None, - ) - - - - - - - - - -class TestNovaTransformationResponse: - """Test response transformation for Nova embeddings.""" - - - - - - - - - -class TestNovaEmbeddingIntegration: - """Integration tests for Nova embeddings through LiteLLM.""" - - - - - - - - -class TestNovaProviderDetection: - """Test provider detection for Nova models.""" - - - - -if __name__ == "__main__": - # Run basic transformation tests - print("Running Nova Embedding Transformation Tests...") - - test_request = TestNovaTransformationRequest() - test_request.test_text_embedding_sync_request() - test_request.test_text_embedding_async_request() - test_request.test_image_embedding_request() - test_request.test_video_embedding_request() - test_request.test_audio_embedding_request() - - test_response = TestNovaTransformationResponse() - test_response.test_text_embedding_response() - test_response.test_multiple_embeddings_response() - test_response.test_async_invoke_response() - - print("All transformation tests passed!") diff --git a/tests/llm_translation/test_cloudflare.py b/tests/llm_translation/test_cloudflare.py deleted file mode 100644 index e8085e091fd..00000000000 --- a/tests/llm_translation/test_cloudflare.py +++ /dev/null @@ -1,99 +0,0 @@ -import asyncio -import json -from unittest.mock import AsyncMock, MagicMock, patch - -import pytest - -from litellm import acompletion, completion -from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler - -FAKE_API_BASE = "https://fake-cloudflare.example.com/client/v4/accounts/fake-acct/ai/v1" -FAKE_API_KEY = "fake-cf-api-key" - - -def _streaming_chunks() -> list[str]: - base = { - "id": "chatcmpl-cf", - "object": "chat.completion.chunk", - "created": 1234567890, - "model": "@cf/meta/llama-2-7b-chat-int8", - } - return [ - json.dumps({**base, "choices": [{"index": 0, "delta": {"content": "I am"}}]}), - json.dumps({**base, "choices": [{"index": 0, "delta": {"content": " a language"}}]}), - json.dumps( - { - **base, - "choices": [ - { - "index": 0, - "delta": {"content": " model."}, - "finish_reason": "stop", - } - ], - } - ), - ] - - -@pytest.mark.parametrize("sync_mode", [False]) -def test_completion_cloudflare_stream(sync_mode): - messages = [{"role": "user", "content": "what llm are you"}] - raw_chunks = _streaming_chunks() - - if sync_mode: - - def _iter_lines(): - for chunk in raw_chunks: - yield f"data: {chunk}" - yield "data: [DONE]" - - mock_resp = MagicMock() - mock_resp.iter_lines.return_value = _iter_lines() - mock_resp.status_code = 200 - mock_resp.headers = {"content-type": "text/event-stream"} - - with patch.object(HTTPHandler, "post", return_value=mock_resp) as mock_post: - response = completion( - model="cloudflare/@cf/meta/llama-2-7b-chat-int8", - messages=messages, - max_tokens=15, - stream=True, - api_base=FAKE_API_BASE, - api_key=FAKE_API_KEY, - ) - chunks_received = list(response) - mock_post.assert_called_once() - else: - - async def _aiter_lines(): - for chunk in raw_chunks: - yield f"data: {chunk}" - yield "data: [DONE]" - - mock_resp = MagicMock() - mock_resp.aiter_lines.return_value = _aiter_lines() - mock_resp.status_code = 200 - mock_resp.headers = {"content-type": "text/event-stream"} - - async def _run(): - with patch.object(AsyncHTTPHandler, "post", new_callable=AsyncMock, return_value=mock_resp) as mock_post: - resp = await acompletion( - model="cloudflare/@cf/meta/llama-2-7b-chat-int8", - messages=messages, - max_tokens=15, - stream=True, - api_base=FAKE_API_BASE, - api_key=FAKE_API_KEY, - ) - received = [] - async for chunk in resp: - received.append(chunk) - mock_post.assert_called_once() - return received - - chunks_received = asyncio.run(_run()) - - assert len(chunks_received) > 0 - content = "".join(c.choices[0].delta.content for c in chunks_received if c.choices[0].delta.content) - assert "language" in content.lower() diff --git a/tests/llm_translation/test_litellm_proxy_provider.py b/tests/llm_translation/test_litellm_proxy_provider.py deleted file mode 100644 index 6e2978c5801..00000000000 --- a/tests/llm_translation/test_litellm_proxy_provider.py +++ /dev/null @@ -1,46 +0,0 @@ -import re - - -import litellm -import pytest - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -def test_litellm_gateway_from_sdk_with_thinking_param(): - with pytest.raises(Exception, match=re.escape("Connection error.")) as exc_info: - response = litellm.completion( - model="litellm_proxy/anthropic.claude-sonnet-4-5-20250929-v1:0", - messages=[{"role": "user", "content": "Hello world"}], - api_base="http://0.0.0.0:4000", - api_key="sk-PIp1h0RekR", - # client=openai_client, - thinking={"type": "enabled", "max_budget": 100}, - ) - e = exc_info.value - assert "Connection error." in str(e) diff --git a/tests/llm_translation/test_llm_response_utils/test_convert_dict_to_chat_completion.py b/tests/llm_translation/test_llm_response_utils/test_convert_dict_to_chat_completion.py deleted file mode 100644 index 74ee2a6dc39..00000000000 --- a/tests/llm_translation/test_llm_response_utils/test_convert_dict_to_chat_completion.py +++ /dev/null @@ -1,93 +0,0 @@ -from datetime import datetime - -from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( - convert_to_model_response_object, -) -from litellm.types.utils import ( - CompletionTokensDetailsWrapper, - Message, - ModelResponse, -) - - -def test_convert_to_model_response_object_function_output(): - """ - Test conversion with function output. - - From here: https://platform.openai.com/docs/api-reference/chat/create - - """ - response_object = { - "id": "chatcmpl-abc123", - "object": "chat.completion", - "created": 1699896916, - "model": "gpt-4o-mini", - "choices": [ - { - "index": 0, - "message": { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_abc123", - "type": "function", - "function": { - "name": "get_current_weather", - "arguments": '{\n"location": "Boston, MA"\n}', - }, - } - ], - }, - "logprobs": None, - "finish_reason": "tool_calls", - } - ], - "usage": { - "prompt_tokens": 82, - "completion_tokens": 17, - "total_tokens": 99, - "completion_tokens_details": {"reasoning_tokens": 0}, - }, - } - - result = convert_to_model_response_object( - model_response_object=ModelResponse(), - response_object=response_object, - stream=False, - start_time=datetime.now(), - end_time=datetime.now(), - hidden_params=None, - _response_headers=None, - convert_tool_call_to_json_mode=False, - ) - - assert isinstance(result, ModelResponse) - assert result.id == "chatcmpl-abc123" - assert result.object == "chat.completion" - assert result.created == 1699896916 - assert result.model == "gpt-4o-mini" - - assert len(result.choices) == 1 - choice = result.choices[0] - assert choice.index == 0 - assert isinstance(choice.message, Message) - assert choice.message.role == "assistant" - assert choice.message.content is None - assert choice.finish_reason == "tool_calls" - - assert len(choice.message.tool_calls) == 1 - tool_call = choice.message.tool_calls[0] - assert tool_call.id == "call_abc123" - assert tool_call.type == "function" - assert tool_call.function.name == "get_current_weather" - assert tool_call.function.arguments == '{\n"location": "Boston, MA"\n}' - - assert result.usage.prompt_tokens == 82 - assert result.usage.completion_tokens == 17 - assert result.usage.total_tokens == 99 - assert result.usage.completion_tokens_details == CompletionTokensDetailsWrapper( - reasoning_tokens=0 - ) - - assert result._hidden_params is not None diff --git a/tests/llm_translation/test_nvidia_nim.py b/tests/llm_translation/test_nvidia_nim.py deleted file mode 100644 index 402035444b1..00000000000 --- a/tests/llm_translation/test_nvidia_nim.py +++ /dev/null @@ -1,72 +0,0 @@ -import json -from datetime import datetime -from typing import Final -from unittest.mock import AsyncMock - -import httpx -import pytest -from openai.types import CreateEmbeddingResponse, Embedding -from openai.types.create_embedding_response import Usage as EmbeddingUsage -from unittest.mock import patch, MagicMock - -import litellm -from litellm import Choices, Message, ModelResponse, EmbeddingResponse, Usage -from litellm import completion -from base_rerank_unit_tests import BaseLLMRerankTest -from tests.capturing_transport import CapturingTransport - -def test_completion_nvidia_nim(): - from openai import OpenAI - - litellm.set_verbose = True - model_name = "nvidia_nim/databricks/dbrx-instruct" - client = OpenAI( - api_key="fake-api-key", - ) - - with patch.object(client.chat.completions.with_raw_response, "create") as mock_client: - try: - completion( - model=model_name, - messages=[ - { - "role": "user", - "content": "What's the weather like in Boston today in Fahrenheit?", - } - ], - presence_penalty=0.5, - frequency_penalty=0.1, - client=client, - ) - except Exception as e: - print(e) - # Add any assertions here to check the response - - mock_client.assert_called_once() - request_body = mock_client.call_args.kwargs - - print("request_body: ", request_body) - - assert request_body["messages"] == [ - { - "role": "user", - "content": "What's the weather like in Boston today in Fahrenheit?", - }, - ] - assert request_body["model"] == "databricks/dbrx-instruct" - assert request_body["frequency_penalty"] == 0.1 - assert request_body["presence_penalty"] == 0.5 - -class TestNvidiaNim(BaseLLMRerankTest): - def get_custom_llm_provider(self) -> litellm.LlmProviders: - return litellm.LlmProviders.NVIDIA_NIM - - def get_base_rerank_call_args(self) -> dict: - return { - "model": "nvidia_nim/nvidia/llama-3_2-nv-rerankqa-1b-v2", - } - - def get_expected_cost(self) -> float: - """Nvidia NIM rerank models are free (cost = 0.0)""" - return 0.0 - diff --git a/tests/llm_translation/test_replicate.py b/tests/llm_translation/test_replicate.py deleted file mode 100644 index 7d08b8bcb57..00000000000 --- a/tests/llm_translation/test_replicate.py +++ /dev/null @@ -1,103 +0,0 @@ -""" -Unit tests for Replicate provider, particularly testing DeepSeek models -""" - -import json -from unittest.mock import AsyncMock, Mock, patch - -import pytest - -import litellm -from litellm.llms.replicate.chat.handler import ( - async_completion, -) -from litellm.llms.replicate.chat.handler import ( - completion as replicate_completion, -) - - -class TestReplicateStartingStatus: - """Test that Replicate handler correctly handles 'starting' status for DeepSeek models""" - - - @patch("litellm.llms.replicate.chat.handler.get_httpx_client") - def test_sync_completion_handles_starting_status(self, mock_get_client): - """Test that sync completion polls correctly when status is 'starting'""" - # Mock the sync HTTP client - mock_client = Mock() - mock_get_client.return_value = mock_client - - # Mock the initial POST response - post_response = Mock() - post_response.json.return_value = { - "id": "test-prediction-id", - "urls": { - "get": "https://api.replicate.com/v1/predictions/test-id", - "cancel": "https://api.replicate.com/v1/predictions/test-id/cancel", - }, - } - mock_client.post.return_value = post_response - - # Mock GET responses - get_response_starting = Mock() - get_response_starting.status_code = 200 - get_response_starting.json.return_value = { - "id": "test-prediction-id", - "status": "starting", - "output": None, - } - - get_response_succeeded = Mock() - get_response_succeeded.status_code = 200 - get_response_succeeded.json.return_value = { - "id": "test-prediction-id", - "status": "succeeded", - "output": ["Hello", " DeepSeek!"], - } - get_response_succeeded.text = json.dumps( - get_response_succeeded.json.return_value - ) - get_response_succeeded.headers = {} - - # Configure mock to return different responses - mock_client.get.side_effect = [get_response_starting, get_response_succeeded] - - # Create mock objects - model_response = litellm.ModelResponse() - model_response.choices = [litellm.Choices()] - model_response.choices[0].message = litellm.Message(content="") - - mock_logging = Mock() - mock_logging.post_call = Mock() - - # Call completion with mock_response to avoid actual API call - with patch("time.sleep"): # Skip sleep delays in test - result = replicate_completion( - model="deepseek-ai/deepseek-v3", - messages=[{"role": "user", "content": "Hi"}], - api_base="https://api.replicate.com", - model_response=model_response, - print_verbose=print, - optional_params={}, - litellm_params={}, - logging_obj=mock_logging, - api_key="test-key", - encoding=None, - headers={}, - ) - - # Assert results - assert result is not None - assert result.choices[0].message.content == "Hello DeepSeek!" - - # Verify GET was called multiple times - assert mock_client.get.call_count >= 1 - - -class TestReplicateOutputFormats: - """Test that Replicate handler handles different output formats from models""" - - - - -# Integration test (requires actual API key - skip in CI) diff --git a/tests/llm_translation/test_triton.py b/tests/llm_translation/test_triton.py index 74598e84c83..b3b9e2b63a3 100644 --- a/tests/llm_translation/test_triton.py +++ b/tests/llm_translation/test_triton.py @@ -15,87 +15,6 @@ from litellm.llms.triton.embedding.transformation import TritonEmbeddingConfig from tests.fake_openai_endpoint import FAKE_OPENAI_API_BASE -@pytest.mark.parametrize("stream", [True, False]) -def test_completion_triton_generate_api(stream): - try: - mock_response = MagicMock() - if stream: - - def mock_iter_lines(): - mock_output = "".join( - [ - 'data: {"model_name":"ensemble","model_version":"1","sequence_end":false,"sequence_id":0,"sequence_start":false,"text_output":"' - + t - + '"}\n\n' - for t in ["I", " am", " an", " AI", " assistant"] - ] - ) - for out in mock_output.split("\n"): - yield out - - mock_response.iter_lines = mock_iter_lines - else: - - def return_val(): - return { - "text_output": "I am an AI assistant", - } - - mock_response.json = return_val - mock_response.status_code = 200 - - with patch( - "litellm.llms.custom_httpx.http_handler.HTTPHandler.post", - return_value=mock_response, - ) as mock_post: - response = litellm.completion( - model="triton/llama-3-8b-instruct", - messages=[{"role": "user", "content": "who are u?"}], - max_tokens=10, - timeout=5, - api_base="http://localhost:8000/generate", - stream=stream, - ) - - # Verify the call was made - mock_post.assert_called_once() - - # Get the arguments passed to the post request - print("call args", mock_post.call_args) - call_kwargs = mock_post.call_args.kwargs # Access kwargs directly - - # Verify URL - if stream: - assert call_kwargs["url"] == "http://localhost:8000/generate_stream" - else: - assert call_kwargs["url"] == "http://localhost:8000/generate" - - # Parse the request data from the JSON string - request_data = json.loads(call_kwargs["data"]) - - # Verify request data - assert request_data["text_input"] == "who are u?" - assert request_data["parameters"]["max_tokens"] == 10 - - # Verify response - if stream: - tokens = ["I", " am", " an", " AI", " assistant", None] - idx = 0 - for chunk in response: - assert chunk.choices[0].delta.content == tokens[idx] - idx += 1 - assert idx == len(tokens) - else: - assert response.choices[0].message.content == "I am an AI assistant" - - except Exception as e: - print("exception", e) - import traceback - - traceback.print_exc() - pytest.fail(f"Error occurred: {e}") - - @pytest.mark.asyncio async def test_triton_embeddings(): try: @@ -111,22 +30,3 @@ async def test_triton_embeddings(): assert response.data[0]["embedding"] == [0.1, 0.2] except Exception as e: pytest.fail(f"Error occurred: {e}") - - -def test_triton_generate_raw_request(): - from litellm.utils import return_raw_request - from litellm.types.utils import CallTypes - - try: - kwargs = { - "model": "triton/llama-3-8b-instruct", - "messages": [{"role": "user", "content": "who are u?"}], - "api_base": "http://localhost:8000/generate", - } - raw_request = return_raw_request(endpoint=CallTypes.completion, kwargs=kwargs) - print("raw_request", raw_request) - assert raw_request is not None - assert "bad_words" not in json.dumps(raw_request["raw_request_body"]) - assert "stop_words" not in json.dumps(raw_request["raw_request_body"]) - except Exception as e: - pytest.fail(f"Error occurred: {e}") diff --git a/tests/llm_translation/test_unit_test_bedrock_invoke.py b/tests/llm_translation/test_unit_test_bedrock_invoke.py deleted file mode 100644 index d237578cd25..00000000000 --- a/tests/llm_translation/test_unit_test_bedrock_invoke.py +++ /dev/null @@ -1,34 +0,0 @@ -import traceback -from dotenv import load_dotenv -import litellm.types -import pytest -from litellm import AmazonInvokeConfig -import json - -load_dotenv() -import io - -from unittest.mock import AsyncMock, Mock, patch - - -# Initialize the transformer -@pytest.fixture -def bedrock_transformer(): - return AmazonInvokeConfig() - - -def test_transform_request_meta_llama(bedrock_transformer): - """Test request transformation for Meta/Llama""" - messages = [{"role": "user", "content": "Hello"}] - - result = bedrock_transformer.transform_request( - model="meta.llama2-70b", - messages=messages, - optional_params={"max_gen_len": 2048}, - litellm_params={}, - headers={}, - ) - - print("transformed request for invoke meta llama=", json.dumps(result, indent=4)) - expected_result = {"prompt": "Hello", "max_gen_len": 2048} - assert result == expected_result diff --git a/tests/llm_translation/test_vcr_classification.py b/tests/llm_translation/test_vcr_classification.py deleted file mode 100644 index d025bc254df..00000000000 --- a/tests/llm_translation/test_vcr_classification.py +++ /dev/null @@ -1,148 +0,0 @@ -"""Unit tests for the VCR classification + observability layer. - -Covers: -- per-item respx detection (module scan, marker, fixture) -- skip-reason tagging in ``apply_vcr_auto_marker_to_items`` -- verdict classification (HIT / MISS:RECORDED / MISS:OVERFLOW / MISS:NOT_PERSISTED / - PARTIAL / NOOP / UNMARKED:LIVE_CALL / UNMARKED:NO_TRAFFIC) -- AWS SigV4 fingerprint stability -- session-end summary rendering -- live-call host classification -""" - -from __future__ import annotations - -import os -import sys -from types import SimpleNamespace -from typing import Optional - -import pytest - -sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..", ".."))) - -from tests._vcr_conftest_common import ( # noqa: E402 - SKIP_REASON_FILE_OPT_OUT, - SKIP_REASON_INCOMPATIBLE, - SKIP_REASON_PRE_MARKED, - SKIP_REASON_RESPX, - SKIP_REASON_RESPX_MODULE, - VCR_SKIP_REASON_USER_ATTR, - VERDICT_HIT, - VERDICT_MISS_NOT_PERSISTED, - VERDICT_MISS_OVERFLOW, - VERDICT_MISS_RECORDED, - VERDICT_NOOP_NO_TRAFFIC, - VERDICT_PARTIAL, - VERDICT_UNMARKED_LIVE_CALL, - VERDICT_UNMARKED_NO_TRAFFIC, - _RESPX_MODULE_CACHE, - _classify_marked_test, - _compute_key_fingerprint, - _is_live_call_host, - _reset_session_stats, - _stable_key_value, - aggregate_report_outcome, - apply_vcr_auto_marker_to_items, - emit_vcr_classification_summary, - install_live_call_probe, - record_vcr_outcome, - session_stats_snapshot, -) - -# --------------------------------------------------------------------------- -# Test doubles -# --------------------------------------------------------------------------- - - -@pytest.fixture -def vcr_enabled(monkeypatch): - monkeypatch.setenv("CASSETTE_REDIS_URL", "redis://stub") - monkeypatch.delenv("LITELLM_VCR_DISABLE", raising=False) - monkeypatch.delenv("PYTEST_XDIST_WORKER", raising=False) - - -@pytest.fixture(autouse=True) -def _reset_module_caches(): - _reset_session_stats() - _RESPX_MODULE_CACHE.clear() - yield - _reset_session_stats() - _RESPX_MODULE_CACHE.clear() - - -# --------------------------------------------------------------------------- -# AWS SigV4 fingerprint stability — the Bedrock cassette overflow root cause -# --------------------------------------------------------------------------- - - -# --------------------------------------------------------------------------- -# Live-call host classification -# --------------------------------------------------------------------------- - - -# --------------------------------------------------------------------------- -# Verdict classification -# --------------------------------------------------------------------------- - - -# --------------------------------------------------------------------------- -# apply_vcr_auto_marker_to_items: skip-reason tagging -# --------------------------------------------------------------------------- - - -# --------------------------------------------------------------------------- -# Session-end summary -# --------------------------------------------------------------------------- - - -# --------------------------------------------------------------------------- -# xdist controller aggregation -# -# _session_stats lives in module-global memory. Under xdist that memory is -# per-worker, so the controller's pytest_terminal_summary would render an -# empty summary without these aggregation hooks. The tests below simulate -# the controller receiving teardown reports produced by workers. -# --------------------------------------------------------------------------- - - -# --------------------------------------------------------------------------- -# Live-call probe -# --------------------------------------------------------------------------- - - -def test_live_call_probe_records_known_llm_hosts(vcr_enabled, monkeypatch): - """The probe should record outbound TCP connections to known LLM - provider hosts (and ignore localhost / RFC1918 / unknown hosts).""" - finalizers = [] - - class _Node: - pass - - request = SimpleNamespace(node=_Node(), addfinalizer=lambda fn: finalizers.append(fn)) - probe = install_live_call_probe(request, None) - assert probe is not None - - import socket - - # Manually invoke the patched function — we don't actually open a - # connection because that would hit the network. The probe records - # at the *call site* before delegating, and the original - # ``socket.create_connection`` will then fail; we swallow that. - try: - socket.create_connection(("api.openai.com", 443), timeout=0.001) - except Exception: - pass - try: - socket.create_connection(("127.0.0.1", 6379), timeout=0.001) - except Exception: - pass - - # Restore via finalizers before asserting so the rest of the test - # session is unaffected. - for fn in finalizers: - fn() - - hosts = getattr(request.node, "vcr_live_call_hosts", []) - assert "api.openai.com" in hosts - assert "127.0.0.1" not in hosts diff --git a/tests/local_testing/test_get_model_info.py b/tests/local_testing/test_get_model_info.py deleted file mode 100644 index 9556dedaa75..00000000000 --- a/tests/local_testing/test_get_model_info.py +++ /dev/null @@ -1,34 +0,0 @@ -import os - -import pytest - -import litellm - - -def test_get_model_info_huggingface_models(monkeypatch): - from litellm import Router - from litellm.types.router import ModelGroupInfo - - monkeypatch.setenv("HUGGINGFACE_API_KEY", "hf_abc123") - - router = Router( - model_list=[ - { - "model_name": "meta-llama/Meta-Llama-3-8B-Instruct", - "litellm_params": { - "model": "huggingface/meta-llama/Meta-Llama-3-8B-Instruct", - "api_base": "https://router.huggingface.co/hf-inference/models/meta-llama/Meta-Llama-3-8B-Instruct", - "api_key": os.environ["HUGGINGFACE_API_KEY"], - }, - } - ] - ) - info = litellm.get_model_info("huggingface/meta-llama/Meta-Llama-3-8B-Instruct") - print("info", info) - assert info is not None - - ModelGroupInfo( - model_group="meta-llama/Meta-Llama-3-8B-Instruct", - providers=["huggingface"], - **info, - ) diff --git a/tests/logging_callback_tests/test_langfuse_dynamic_credentials.py b/tests/logging_callback_tests/test_langfuse_dynamic_credentials.py deleted file mode 100644 index b6cac87631c..00000000000 --- a/tests/logging_callback_tests/test_langfuse_dynamic_credentials.py +++ /dev/null @@ -1,32 +0,0 @@ -import litellm - - - - - - - - -def test_upstream_langfuse_env_only_warns_and_opens_no_second_channel(monkeypatch, caplog): - """UPSTREAM_LANGFUSE_* configured a second v2 ingestion client. v4 has one export channel per - credential set, so the values are ignored with a startup warning and never build anything.""" - from litellm.integrations.langfuse import langfuse_sdk - from litellm.integrations.langfuse.langfuse import LangFuseLogger - - monkeypatch.setattr(litellm, "initialized_langfuse_clients", 0) - monkeypatch.setattr(langfuse_sdk, "_TRACING", {}) - monkeypatch.setenv("LANGFUSE_MOCK", "true") - monkeypatch.setenv("UPSTREAM_LANGFUSE_SECRET_KEY", "upstream-secret") - monkeypatch.setenv("UPSTREAM_LANGFUSE_PUBLIC_KEY", "upstream-public") - monkeypatch.setenv("UPSTREAM_LANGFUSE_HOST", "https://upstream.example") - - with caplog.at_level("WARNING", logger="LiteLLM"): - logger = LangFuseLogger( - langfuse_public_key="public", - langfuse_secret="secret", - langfuse_host="https://langfuse.example", - ) - - assert any("UPSTREAM_LANGFUSE_* is no longer supported" in record.getMessage() for record in caplog.records) - assert [lease.tracing for lease in langfuse_sdk._TRACING.values()] == [logger.tracing] - assert all(key.public_key == "public" for key in langfuse_sdk._TRACING) diff --git a/tests/logging_callback_tests/test_langfuse_unit_tests.py b/tests/logging_callback_tests/test_langfuse_unit_tests.py index b7554cb129b..1a834a37c46 100644 --- a/tests/logging_callback_tests/test_langfuse_unit_tests.py +++ b/tests/logging_callback_tests/test_langfuse_unit_tests.py @@ -1,308 +1,4 @@ -import os - - -import pytest -from litellm.integrations.langfuse.langfuse import ( - LangFuseLogger, -) -from litellm.integrations.langfuse.langfuse_handler import LangFuseHandler from litellm.litellm_core_utils.litellm_logging import DynamicLoggingCache -from unittest.mock import Mock, patch -from litellm.types.utils import ( - StandardLoggingPayload, - StandardLoggingModelInformation, - StandardLoggingMetadata, - StandardLoggingHiddenParams, - StandardCallbackDynamicParams, - ModelResponse, - Choices, - Message, - TextCompletionResponse, - TextChoices, -) - - -def create_standard_logging_payload() -> StandardLoggingPayload: - return StandardLoggingPayload( - id="test_id", - call_type="completion", - response_cost=0.1, - response_cost_failure_debug_info=None, - status="success", - total_tokens=30, - prompt_tokens=20, - completion_tokens=10, - startTime=1234567890.0, - endTime=1234567891.0, - completionStartTime=1234567890.5, - model_map_information=StandardLoggingModelInformation( - model_map_key="gpt-5-mini", model_map_value=None - ), - model="gpt-5-mini", - model_id="model-123", - model_group="openai-gpt", - api_base="https://api.openai.com", - metadata=StandardLoggingMetadata( - user_api_key_hash="test_hash", - user_api_key_org_id=None, - user_api_key_alias="test_alias", - user_api_key_team_id="test_team", - user_api_key_user_id="test_user", - user_api_key_team_alias="test_team_alias", - spend_logs_metadata=None, - requester_ip_address="127.0.0.1", - requester_metadata=None, - ), - cache_hit=False, - cache_key=None, - saved_cache_cost=0.0, - request_tags=[], - end_user=None, - requester_ip_address="127.0.0.1", - messages=[{"role": "user", "content": "Hello, world!"}], - response={"choices": [{"message": {"content": "Hi there!"}}]}, - error_str=None, - model_parameters={"stream": True}, - hidden_params=StandardLoggingHiddenParams( - model_id="model-123", - cache_key=None, - api_base="https://api.openai.com", - response_cost="0.1", - additional_headers=None, - ), - ) - - -@pytest.fixture -def dynamic_logging_cache(): - return DynamicLoggingCache() - - -global_langfuse_logger = LangFuseLogger( - langfuse_public_key="global_public_key", - langfuse_secret="global_secret", - langfuse_host="https://global.langfuse.com", -) - - -# IMPORTANT: Test that passing both langfuse_secret_key and langfuse_secret works -standard_params_1 = StandardCallbackDynamicParams( - langfuse_public_key="test_public_key", - langfuse_secret="test_secret", - langfuse_host="https://test.langfuse.com", -) - -standard_params_2 = StandardCallbackDynamicParams( - langfuse_public_key="test_public_key", - langfuse_secret_key="test_secret", - langfuse_host="https://test.langfuse.com", -) - - -@pytest.mark.parametrize("globalLangfuseLogger", [None, global_langfuse_logger]) -@pytest.mark.parametrize("standard_params", [standard_params_1, standard_params_2]) -def test_get_langfuse_logger_for_request_with_dynamic_params( - dynamic_logging_cache, globalLangfuseLogger, standard_params -): - """ - If StandardCallbackDynamicParams contain langfuse credentials the returned Langfuse logger should use the dynamic params - - the new Langfuse logger should be cached - - Even if globalLangfuseLogger is provided, it should use dynamic params if they are passed - """ - - result = LangFuseHandler.get_langfuse_logger_for_request( - standard_callback_dynamic_params=standard_params, - in_memory_dynamic_logger_cache=dynamic_logging_cache, - globalLangfuseLogger=globalLangfuseLogger, - ) - - assert isinstance(result, LangFuseLogger) - assert result.public_key == "test_public_key" - assert result.secret_key == "test_secret" - assert result.langfuse_host == "https://test.langfuse.com" - - logger_for_identical_repeat_request = LangFuseHandler.get_langfuse_logger_for_request( - standard_callback_dynamic_params=standard_params, - in_memory_dynamic_logger_cache=dynamic_logging_cache, - globalLangfuseLogger=globalLangfuseLogger, - ) - assert logger_for_identical_repeat_request is result - - -@pytest.mark.parametrize("globalLangfuseLogger", [None, global_langfuse_logger]) -def test_get_langfuse_logger_for_request_with_no_dynamic_params( - dynamic_logging_cache, globalLangfuseLogger -): - """ - If StandardCallbackDynamicParams are not provided, the globalLangfuseLogger should be returned - """ - result = LangFuseHandler.get_langfuse_logger_for_request( - standard_callback_dynamic_params=StandardCallbackDynamicParams(), - in_memory_dynamic_logger_cache=dynamic_logging_cache, - globalLangfuseLogger=globalLangfuseLogger, - ) - - assert result is not None - assert isinstance(result, LangFuseLogger) - - if globalLangfuseLogger is not None: - assert result.public_key == "global_public_key" - assert result.secret_key == "global_secret" - assert result.langfuse_host == "https://global.langfuse.com" - - -def test_dynamic_langfuse_credentials_are_passed(): - # Test when credentials are passed - params_with_credentials = StandardCallbackDynamicParams( - langfuse_public_key="test_key", - langfuse_secret="test_secret", - langfuse_host="https://test.langfuse.com", - ) - assert ( - LangFuseHandler._dynamic_langfuse_credentials_are_passed( - params_with_credentials - ) - is True - ) - - # Test when no credentials are passed - params_without_credentials = StandardCallbackDynamicParams() - assert ( - LangFuseHandler._dynamic_langfuse_credentials_are_passed( - params_without_credentials - ) - is False - ) - - # Test when only some credentials are passed - params_partial_credentials = StandardCallbackDynamicParams( - langfuse_public_key="test_key" - ) - assert ( - LangFuseHandler._dynamic_langfuse_credentials_are_passed( - params_partial_credentials - ) - is True - ) - - -def test_get_dynamic_langfuse_logging_config(): - # Test with dynamic params - dynamic_params = StandardCallbackDynamicParams( - langfuse_public_key="dynamic_key", - langfuse_secret="dynamic_secret", - langfuse_host="https://dynamic.langfuse.com", - ) - config = LangFuseHandler.get_dynamic_langfuse_logging_config(dynamic_params) - assert config["langfuse_public_key"] == "dynamic_key" - assert config["langfuse_secret"] == "dynamic_secret" - assert config["langfuse_host"] == "https://dynamic.langfuse.com" - - # Test with no dynamic params - empty_params = StandardCallbackDynamicParams() - config = LangFuseHandler.get_dynamic_langfuse_logging_config(empty_params) - assert config["langfuse_public_key"] is None - assert config["langfuse_secret"] is None - assert config["langfuse_host"] is None - - -def test_return_global_langfuse_logger(): - mock_cache = Mock() - global_logger = LangFuseLogger( - langfuse_public_key="global_key", langfuse_secret="global_secret" - ) - - # Test with existing global logger - result = LangFuseHandler._return_global_langfuse_logger(global_logger, mock_cache) - assert result == global_logger - - # Test without global logger, but with cached logger, should return cached logger - mock_cache.get_cache.return_value = global_logger - result = LangFuseHandler._return_global_langfuse_logger(None, mock_cache) - assert result == global_logger - - # Test without global logger and without cached logger, should create new logger - mock_cache.get_cache.return_value = None - with patch.object( - LangFuseHandler, - "_create_langfuse_logger_from_credentials", - return_value=global_logger, - ): - result = LangFuseHandler._return_global_langfuse_logger(None, mock_cache) - assert result == global_logger - - -def test_get_langfuse_logger_for_request_with_cached_logger(): - """ - Test that get_langfuse_logger_for_request returns the cached logger if it exists when dynamic params are passed - """ - mock_cache = Mock() - cached_logger = LangFuseLogger( - langfuse_public_key="cached_key", langfuse_secret="cached_secret" - ) - mock_cache.get_cache.return_value = cached_logger - - dynamic_params = StandardCallbackDynamicParams( - langfuse_public_key="test_key", - langfuse_secret="test_secret", - langfuse_host="https://test.langfuse.com", - ) - - result = LangFuseHandler.get_langfuse_logger_for_request( - standard_callback_dynamic_params=dynamic_params, - in_memory_dynamic_logger_cache=mock_cache, - globalLangfuseLogger=None, - ) - - assert result == cached_logger - mock_cache.get_cache.assert_called_once() - - -def test_get_langfuse_tags(): - """ - Test that _get_langfuse_tags correctly extracts tags from the standard logging payload - """ - # Create a mock logging payload with tags - mock_payload = create_standard_logging_payload() - mock_payload["request_tags"] = ["tag1", "tag2", "test_tag"] - - # Test with payload containing tags - result = global_langfuse_logger._get_langfuse_tags(mock_payload) - assert result == ["tag1", "tag2", "test_tag"] - - # Test with payload without tags - mock_payload["request_tags"] = None - result = global_langfuse_logger._get_langfuse_tags(mock_payload) - assert result == [] - - # Test with empty tags list - mock_payload["request_tags"] = [] - result = global_langfuse_logger._get_langfuse_tags(mock_payload) - assert result == [] - - -@patch.dict(os.environ, {}, clear=True) # Start with empty environment -def test_get_langfuse_flush_interval(): - """ - Test that _get_langfuse_flush_interval correctly reads from environment variable - or falls back to the provided flush_interval - """ - default_interval = 60 - - # Test when env var is not set - result = LangFuseLogger._get_langfuse_flush_interval( - flush_interval=default_interval - ) - assert result == default_interval - - # Test when env var is set - with patch.dict(os.environ, {"LANGFUSE_FLUSH_INTERVAL": "120"}): - result = LangFuseLogger._get_langfuse_flush_interval( - flush_interval=default_interval - ) - assert result == 120 def test_langfuse_e2e_sync(monkeypatch): @@ -363,246 +59,3 @@ def test_langfuse_e2e_sync(monkeypatch): assert received_paths, "langfuse exported nothing" assert all(path.endswith("/api/public/otel/v1/traces") for path in received_paths) - - -def test_get_chat_content_for_langfuse(): - """ - Test that _get_chat_content_for_langfuse correctly extracts content from chat completion responses - """ - # Test with valid response - mock_response = ModelResponse( - choices=[Choices(message=Message(role="assistant", content="Hello world"))] - ) - - result = LangFuseLogger._get_chat_content_for_langfuse(mock_response) - assert result["content"] == "Hello world" - assert result["role"] == "assistant" - - # Test with empty choices - mock_response = ModelResponse(choices=[]) - result = LangFuseLogger._get_chat_content_for_langfuse(mock_response) - assert result is None - - -def test_get_text_completion_content_for_langfuse(): - """ - Test that _get_text_completion_content_for_langfuse correctly extracts content from text completion responses - """ - # Test with valid response - mock_response = TextCompletionResponse(choices=[TextChoices(text="Hello world")]) - result = LangFuseLogger._get_text_completion_content_for_langfuse(mock_response) - assert result == "Hello world" - - # Test with empty choices - mock_response = TextCompletionResponse(choices=[]) - result = LangFuseLogger._get_text_completion_content_for_langfuse(mock_response) - assert result is None - - # Test with no choices field - mock_response = TextCompletionResponse() - result = LangFuseLogger._get_text_completion_content_for_langfuse(mock_response) - assert result is None - - -def test_apply_masking_function_with_string(): - """ - Test that _apply_masking_function correctly applies masking to strings - """ - import re - - def mask_credit_cards(data): - if isinstance(data, str): - return re.sub(r"\b\d{4}[\s-]?\d{4}[\s-]?\d{4}[\s-]?\d{4}\b", "[CARD]", data) - return data - - # Test with string containing credit card - input_str = "My card is 4532-1234-5678-9012" - result = LangFuseLogger._apply_masking_function(input_str, mask_credit_cards) - assert result == "My card is [CARD]" - assert "4532" not in result - - # Test with string without sensitive data - input_str = "Hello world" - result = LangFuseLogger._apply_masking_function(input_str, mask_credit_cards) - assert result == "Hello world" - - -def test_apply_masking_function_with_dict(): - """ - Test that _apply_masking_function correctly applies masking to nested dicts - """ - import re - - def mask_emails(data): - if isinstance(data, str): - return re.sub(r"[\w\.-]+@[\w\.-]+", "[EMAIL]", data) - return data - - # Test with dict containing messages - input_dict = { - "messages": [{"role": "user", "content": "My email is test@example.com"}] - } - result = LangFuseLogger._apply_masking_function(input_dict, mask_emails) - assert result["messages"][0]["content"] == "My email is [EMAIL]" - assert "test@example.com" not in str(result) - - -def test_apply_masking_function_with_none(): - """ - Test that _apply_masking_function handles None correctly - """ - - def dummy_mask(data): - return data - - result = LangFuseLogger._apply_masking_function(None, dummy_mask) - assert result is None - - -def test_apply_masking_function_with_list(): - """ - Test that _apply_masking_function correctly applies masking to lists - """ - import re - - def mask_ssn(data): - if isinstance(data, str): - return re.sub(r"\b\d{3}-\d{2}-\d{4}\b", "[SSN]", data) - return data - - input_list = ["SSN: 123-45-6789", "No sensitive data here"] - result = LangFuseLogger._apply_masking_function(input_list, mask_ssn) - assert result[0] == "SSN: [SSN]" - assert result[1] == "No sensitive data here" - - -def test_masking_function_isolated_from_other_loggers(): - """ - Test that langfuse_masking_function is extracted from metadata and stored separately. - This ensures the callable doesn't leak to other logging integrations. - """ - from litellm.litellm_core_utils.litellm_logging import ( - scrub_sensitive_keys_in_metadata, - ) - - def my_masking_fn(data): - return data - - # Simulate litellm_params with masking function in metadata - litellm_params = { - "metadata": { - "langfuse_masking_function": my_masking_fn, - "other_key": "other_value", - } - } - - # Scrub should extract the function - result = scrub_sensitive_keys_in_metadata(litellm_params) - - # Function should be removed from metadata (won't leak to other loggers) - assert "langfuse_masking_function" not in result["metadata"] - - # Function should be stored in dedicated key for Langfuse to access - assert result.get("_langfuse_masking_function") == my_masking_fn - - # Other metadata should remain intact - assert result["metadata"]["other_key"] == "other_value" - - -def test_masking_function_not_in_metadata_when_not_provided(): - """ - Test that scrub_sensitive_keys_in_metadata works normally when no masking function is provided. - """ - from litellm.litellm_core_utils.litellm_logging import ( - scrub_sensitive_keys_in_metadata, - ) - - litellm_params = { - "metadata": { - "some_key": "some_value", - } - } - - result = scrub_sensitive_keys_in_metadata(litellm_params) - - # No _langfuse_masking_function should be added - assert "_langfuse_masking_function" not in result - - # Original metadata should be unchanged - assert result["metadata"]["some_key"] == "some_value" - - -def test_langfuse_model_parameters_no_secret_leakage(): - """ - Test that sensitive keys in optional_params (api_key, secret_fields, - authorization headers, etc.) are NOT passed to Langfuse as modelParameters. - Only whitelisted model parameters (temperature, top_p, etc.) should survive. - """ - from litellm.litellm_core_utils.model_param_helper import ModelParamHelper - - optional_params_with_secrets = { - # Safe params that should be kept - "temperature": 0.7, - "top_p": 0.9, - "max_tokens": 100, - "stream": True, - # Sensitive params that must NOT leak - "api_key": "sk-secret-key-12345", - "api_base": "https://my-private-endpoint.com", - "secret_fields": {"raw_headers": {"Authorization": "Bearer sk-super-secret"}}, - "authorization": "Bearer sk-another-secret", - "headers": {"X-Api-Key": "secret-header-value"}, - } - - sanitized = ModelParamHelper.get_standard_logging_model_parameters( - optional_params_with_secrets - ) - - # Safe params should be present - assert sanitized["temperature"] == 0.7 - assert sanitized["top_p"] == 0.9 - assert sanitized["max_tokens"] == 100 - assert sanitized["stream"] is True - - # Sensitive params must be excluded - assert "api_key" not in sanitized - assert "api_base" not in sanitized - assert "secret_fields" not in sanitized - assert "authorization" not in sanitized - assert "headers" not in sanitized - - -def test_langfuse_v2_uses_standard_logging_model_parameters(): - """ - Test that _log_langfuse_v2 uses sanitized model_parameters from - standard_logging_object instead of raw optional_params, preventing - secret leakage to Langfuse traces. - """ - standard_logging_object = create_standard_logging_payload() - # Simulate standard_logging_object having safe model_parameters - standard_logging_object["model_parameters"] = {"temperature": 0.5, "stream": True} - - # optional_params has secrets — these should NOT be used - optional_params_with_secrets = { - "temperature": 0.5, - "api_key": "sk-secret-key-12345", - "secret_fields": {"raw_headers": {"Authorization": "Bearer sk-secret"}}, - } - - # When standard_logging_object is available, its model_parameters should be used - sanitized = standard_logging_object.get( - "model_parameters", optional_params_with_secrets - ) - assert "api_key" not in sanitized - assert "secret_fields" not in sanitized - assert sanitized["temperature"] == 0.5 - - # When standard_logging_object is None, ModelParamHelper should filter - from litellm.litellm_core_utils.model_param_helper import ModelParamHelper - - fallback_sanitized = ModelParamHelper.get_standard_logging_model_parameters( - optional_params_with_secrets - ) - assert "api_key" not in fallback_sanitized - assert "secret_fields" not in fallback_sanitized - assert fallback_sanitized["temperature"] == 0.5 diff --git a/tests/logging_callback_tests/test_sqs_logger.py b/tests/logging_callback_tests/test_sqs_logger.py index 31e9ffc5517..f7b8b5e692f 100644 --- a/tests/logging_callback_tests/test_sqs_logger.py +++ b/tests/logging_callback_tests/test_sqs_logger.py @@ -145,142 +145,6 @@ async def test_async_sqs_logger_error_flush(): assert payload_data["messages"][0]["content"] == "hello" -# ============================================================================= -# 📥 Logging Queue Tests -# ============================================================================= - - -# ============================================================================= -# 🧾 async_send_batch Tests -# ============================================================================= - - -@pytest.mark.asyncio -async def test_async_send_batch_does_not_await_send_directly(monkeypatch): - # create_task stays real here: with it mocked out the await_count assertion - # below would hold trivially. Every task it spawns is cancelled at the end, - # including the infinite periodic_flush the SQSLogger constructor starts. - monkeypatch.setattr("litellm.aws_sqs_callback_params", {}) - spawned = [] - real_create_task = asyncio.create_task - - def spy_create_task(coro, *args, **kwargs): - task = real_create_task(coro, *args, **kwargs) - spawned.append(task) - return task - - monkeypatch.setattr(asyncio, "create_task", spy_create_task) - - logger = SQSLogger(sqs_queue_url="https://example.com", sqs_region_name="us-west-2") - logger.async_send_message = AsyncMock() - logger.log_queue = [{"log": 1}, {"log": 2}] - - try: - await logger.async_send_batch() - assert logger.async_send_message.await_count == 0 - finally: - for task in spawned: - task.cancel() - await asyncio.gather(*spawned, return_exceptions=True) - - -# ============================================================================= -# 🔐 AppCrypto Tests -# ============================================================================= - - -def test_appcrypto_encrypt_decrypt_roundtrip(): - key = os.urandom(32) - crypto = AppCrypto(key) - data = {"event": "test", "value": 42} - aad = b"context" - enc = crypto.encrypt_json(data, aad=aad) - dec = crypto.decrypt_json(enc, aad=aad) - assert dec == data - - -def test_appcrypto_invalid_key_length(): - with pytest.raises(ValueError, match="32 bytes"): - AppCrypto(b"short") - - -# ============================================================================= -# 🪣 SQSLogger Initialization Tests -# ============================================================================= - - -def test_sqs_logger_init_without_encryption(monkeypatch): - monkeypatch.setattr("litellm.aws_sqs_callback_params", {}) - # Patch asyncio.create_task to avoid RuntimeError - monkeypatch.setattr(asyncio, "create_task", MagicMock()) - logger = SQSLogger(sqs_queue_url="https://example.com", sqs_region_name="us-west-2") - assert logger.sqs_queue_url == "https://example.com" - assert logger.app_crypto is None - - -def test_sqs_logger_init_with_encryption(monkeypatch): - monkeypatch.setattr("litellm.aws_sqs_callback_params", {}) - monkeypatch.setattr(asyncio, "create_task", MagicMock()) - key_b64 = base64.b64encode(os.urandom(32)).decode() - - logger = SQSLogger( - sqs_queue_url="https://example.com", - sqs_region_name="us-west-2", - sqs_aws_use_application_level_encryption=True, - sqs_app_encryption_key_b64=key_b64, - sqs_app_encryption_aad="tenant=bill", - ) - assert logger.app_crypto is not None - assert logger.sqs_app_encryption_aad == "tenant=bill" - - -def test_sqs_logger_init_with_encryption_missing_key(monkeypatch): - monkeypatch.setattr("litellm.aws_sqs_callback_params", {}) - monkeypatch.setattr(asyncio, "create_task", MagicMock()) - with pytest.raises(ValueError, match="required when encryption is enabled"): - SQSLogger( - sqs_queue_url="https://example.com", - sqs_region_name="us-west-2", - sqs_aws_use_application_level_encryption=True, - ) - - -# ============================================================================= -# 📥 Logging Queue Tests -# ============================================================================= - - -@pytest.mark.asyncio -async def test_async_log_success_event_adds_to_queue(monkeypatch): - monkeypatch.setattr("litellm.aws_sqs_callback_params", {}) - monkeypatch.setattr(asyncio, "create_task", MagicMock()) - logger = SQSLogger(sqs_queue_url="https://example.com", sqs_region_name="us-west-2") - - fake_payload = {"some": "data"} - await logger.async_log_success_event( - {"standard_logging_object": fake_payload}, None, None, None - ) - assert fake_payload in logger.log_queue - - -@pytest.mark.asyncio -async def test_async_log_failure_event_adds_to_queue(monkeypatch): - monkeypatch.setattr("litellm.aws_sqs_callback_params", {}) - monkeypatch.setattr(asyncio, "create_task", MagicMock()) - logger = SQSLogger(sqs_queue_url="https://example.com", sqs_region_name="us-west-2") - - fake_payload = {"fail": True} - await logger.async_log_failure_event( - {"standard_logging_object": fake_payload}, None, None, None - ) - assert fake_payload in logger.log_queue - - -# ============================================================================= -# 🧾 async_send_batch Tests -# ============================================================================= - - @pytest.mark.asyncio async def test_async_send_batch_triggers_tasks(monkeypatch): monkeypatch.setattr("litellm.aws_sqs_callback_params", {}) @@ -295,55 +159,6 @@ async def test_async_send_batch_triggers_tasks(monkeypatch): asyncio.create_task.assert_called() -@pytest.mark.asyncio -async def test_strip_base64_removes_file_and_nontext_entries(): - logger = SQSLogger(sqs_strip_base64_files=True) - - payload = { - "messages": [ - { - "role": "user", - "content": [ - {"type": "text", "text": "Hello world"}, - { - "type": "image", - "file": {"file_data": "data:image/png;base64,AAAA"}, - }, - { - "type": "file", - "file": {"file_data": "data:application/pdf;base64,BBBB"}, - }, - ], - }, - { - "role": "assistant", - "content": [ - {"type": "text", "text": "Response"}, - { - "type": "audio", - "file": {"file_data": "data:audio/wav;base64,CCCC"}, - }, - ], - }, - ] - } - - stripped = await logger._strip_base64_from_messages(payload) - - # 1️⃣ All file/image/audio entries removed - assert len(stripped["messages"][0]["content"]) == 1 - assert stripped["messages"][0]["content"][0]["text"] == "Hello world" - - assert len(stripped["messages"][1]["content"]) == 1 - assert stripped["messages"][1]["content"][0]["text"] == "Response" - - # 2️⃣ No residual 'file' keys left - for msg in stripped["messages"]: - for content in msg["content"]: - assert "file" not in content - assert content.get("type") == "text" - - @pytest.mark.asyncio async def test_strip_base64_keeps_non_file_content(): logger = SQSLogger(sqs_strip_base64_files=True) @@ -364,109 +179,3 @@ async def test_strip_base64_keeps_non_file_content(): # Should not modify normal text messages assert stripped["messages"][0]["content"] == payload["messages"][0]["content"] - - -@pytest.mark.asyncio -async def test_strip_base64_handles_empty_or_missing_messages(): - logger = SQSLogger(sqs_strip_base64_files=True) - - payload_no_messages = {} - stripped1 = await logger._strip_base64_from_messages(payload_no_messages) - assert stripped1 == payload_no_messages - - payload_empty = {"messages": []} - stripped2 = await logger._strip_base64_from_messages(payload_empty) - assert stripped2 == payload_empty - - -@pytest.mark.asyncio -async def test_strip_base64_mixed_nested_objects(): - """ - Handles weird/nested content structures gracefully. - """ - logger = SQSLogger(sqs_strip_base64_files=True) - - payload = { - "messages": [ - { - "role": "system", - "content": [ - {"type": "text", "text": "Keep me"}, - {"type": "custom", "metadata": "ignore but non-text"}, - {"foo": "bar"}, - {"file": {"file_data": "data:application/pdf;base64,XXX"}}, - ], - "extra": {"trace_id": "123"}, - } - ] - } - - stripped = await logger._strip_base64_from_messages(payload) - - # 'custom' (non-text) and 'file' entries removed - content = stripped["messages"][0]["content"] - assert len(content) == 2 - assert {"type": "text", "text": "Keep me"} in content - assert {"foo": "bar"} in content - # Other metadata stays - assert stripped["messages"][0]["extra"]["trace_id"] == "123" - - -@pytest.mark.asyncio -async def test_strip_base64_recursive_redaction(): - logger = SQSLogger(sqs_strip_base64_files=True) - payload = { - "messages": [ - { - "content": [ - {"type": "text", "text": "normal text"}, - { - "type": "text", - "text": "data:image/png;base64,iVBORw0KGgoAAAANSUhEUg", - }, - { - "type": "text", - "text": "Nested: {'data': 'data:application/pdf;base64,AAA...'}", - }, - {"file": {"file_data": "data:application/pdf;base64,AAAA"}}, - {"metadata": {"preview": "data:audio/mp3;base64,AAAAA=="}}, - ] - } - ] - } - - result = await logger._strip_base64_from_messages(payload) - content = result["messages"][0]["content"] - - # Dropped file-type entry - assert not any("file" in c for c in content) - # Base64 redacted globally - for c in content: - if isinstance(c, dict): - s = json.dumps(c).lower() - # allow "[base64_redacted]" but nothing else - assert "base64," not in s, f"Found real base64 blob in: {s}" - - -@pytest.mark.asyncio -async def test_async_health_check_healthy(monkeypatch): - monkeypatch.setattr("litellm.aws_sqs_callback_params", {}) - monkeypatch.setattr(asyncio, "create_task", MagicMock()) - logger = SQSLogger(sqs_queue_url="https://example.com", sqs_region_name="us-west-2") - logger.async_send_message = AsyncMock(return_value=None) - - result = await logger.async_health_check() - assert result["status"] == "healthy" - assert result.get("error_message") is None - - -@pytest.mark.asyncio -async def test_async_health_check_unhealthy(monkeypatch): - monkeypatch.setattr("litellm.aws_sqs_callback_params", {}) - monkeypatch.setattr(asyncio, "create_task", MagicMock()) - logger = SQSLogger(sqs_queue_url="https://example.com", sqs_region_name="us-west-2") - logger.async_send_message = AsyncMock(side_effect=Exception("boom")) - - result = await logger.async_health_check() - assert result["status"] == "unhealthy" - assert "boom" in (result.get("error_message") or "") diff --git a/tests/unit/integrations/langfuse/test_langfuse_dynamic_credentials.py b/tests/unit/integrations/langfuse/test_langfuse_dynamic_credentials.py index 8a0cc91b546..fc2d20169e4 100644 --- a/tests/unit/integrations/langfuse/test_langfuse_dynamic_credentials.py +++ b/tests/unit/integrations/langfuse/test_langfuse_dynamic_credentials.py @@ -104,3 +104,23 @@ def test_langfuse_handler_accepts_secret_key_alias(monkeypatch): assert captured["allow_env_credentials"] is False assert captured["cached_service_name"] == "langfuse" assert captured["cached_logging_obj"] is logger + + +def test_upstream_langfuse_env_only_warns_and_opens_no_second_channel(monkeypatch, caplog): + monkeypatch.setattr(litellm, "initialized_langfuse_clients", 0) + monkeypatch.setattr(langfuse_sdk, "_TRACING", {}) + monkeypatch.setenv("LANGFUSE_MOCK", "true") + monkeypatch.setenv("UPSTREAM_LANGFUSE_SECRET_KEY", "upstream-secret") + monkeypatch.setenv("UPSTREAM_LANGFUSE_PUBLIC_KEY", "upstream-public") + monkeypatch.setenv("UPSTREAM_LANGFUSE_HOST", "https://upstream.example") + + with caplog.at_level("WARNING", logger="LiteLLM"): + logger = LangFuseLogger( + langfuse_public_key="public", + langfuse_secret="secret", + langfuse_host="https://langfuse.example", + ) + + assert any("UPSTREAM_LANGFUSE_* is no longer supported" in record.getMessage() for record in caplog.records) + assert [lease.tracing for lease in langfuse_sdk._TRACING.values()] == [logger.tracing] + assert all(key.public_key == "public" for key in langfuse_sdk._TRACING) diff --git a/tests/unit/integrations/langfuse/test_langfuse_handler.py b/tests/unit/integrations/langfuse/test_langfuse_handler.py new file mode 100644 index 00000000000..55e2799ccf4 --- /dev/null +++ b/tests/unit/integrations/langfuse/test_langfuse_handler.py @@ -0,0 +1,163 @@ +from collections.abc import Iterator +from unittest.mock import Mock, patch + +import httpx +import pytest +import respx + +from litellm.integrations.langfuse.langfuse import LangFuseLogger +from litellm.integrations.langfuse.langfuse_handler import LangFuseHandler +from litellm.litellm_core_utils.specialty_caches.dynamic_logging_cache import DynamicLoggingCache +from litellm.types.utils import StandardCallbackDynamicParams + + +@pytest.fixture(autouse=True) +def langfuse_api_unreachable() -> Iterator[respx.MockRouter]: + with respx.mock(assert_all_called=False) as router: + router.route().mock(side_effect=httpx.ConnectError("langfuse is unreachable")) + yield router + + +@pytest.fixture +def dynamic_logging_cache(): + return DynamicLoggingCache() + + +@pytest.fixture +def globalLangfuseLogger(request: pytest.FixtureRequest) -> LangFuseLogger | None: + if request.param is None: + return None + return LangFuseLogger( + langfuse_public_key="global_public_key", + langfuse_secret="global_secret", + langfuse_host="https://global.langfuse.com", + ) + + +standard_params_1 = StandardCallbackDynamicParams( + langfuse_public_key="test_public_key", + langfuse_secret="test_secret", + langfuse_host="https://test.langfuse.com", +) + +standard_params_2 = StandardCallbackDynamicParams( + langfuse_public_key="test_public_key", + langfuse_secret_key="test_secret", + langfuse_host="https://test.langfuse.com", +) + + +@pytest.mark.parametrize("globalLangfuseLogger", [None, "global"], indirect=True) +@pytest.mark.parametrize("standard_params", [standard_params_1, standard_params_2]) +def test_get_langfuse_logger_for_request_with_dynamic_params( + dynamic_logging_cache, globalLangfuseLogger, standard_params +): + result = LangFuseHandler.get_langfuse_logger_for_request( + standard_callback_dynamic_params=standard_params, + in_memory_dynamic_logger_cache=dynamic_logging_cache, + globalLangfuseLogger=globalLangfuseLogger, + ) + + assert isinstance(result, LangFuseLogger) + assert result.public_key == "test_public_key" + assert result.secret_key == "test_secret" + assert result.langfuse_host == "https://test.langfuse.com" + + logger_for_identical_repeat_request = LangFuseHandler.get_langfuse_logger_for_request( + standard_callback_dynamic_params=standard_params, + in_memory_dynamic_logger_cache=dynamic_logging_cache, + globalLangfuseLogger=globalLangfuseLogger, + ) + assert logger_for_identical_repeat_request is result + + +@pytest.mark.parametrize("globalLangfuseLogger", [None, "global"], indirect=True) +def test_get_langfuse_logger_for_request_with_no_dynamic_params(dynamic_logging_cache, globalLangfuseLogger): + result = LangFuseHandler.get_langfuse_logger_for_request( + standard_callback_dynamic_params=StandardCallbackDynamicParams(), + in_memory_dynamic_logger_cache=dynamic_logging_cache, + globalLangfuseLogger=globalLangfuseLogger, + ) + + assert result is not None + assert isinstance(result, LangFuseLogger) + + if globalLangfuseLogger is not None: + assert result.public_key == "global_public_key" + assert result.secret_key == "global_secret" + assert result.langfuse_host == "https://global.langfuse.com" + + +def test_dynamic_langfuse_credentials_are_passed(): + params_with_credentials = StandardCallbackDynamicParams( + langfuse_public_key="test_key", + langfuse_secret="test_secret", + langfuse_host="https://test.langfuse.com", + ) + assert LangFuseHandler._dynamic_langfuse_credentials_are_passed(params_with_credentials) is True + + params_without_credentials = StandardCallbackDynamicParams() + assert LangFuseHandler._dynamic_langfuse_credentials_are_passed(params_without_credentials) is False + + params_partial_credentials = StandardCallbackDynamicParams(langfuse_public_key="test_key") + assert LangFuseHandler._dynamic_langfuse_credentials_are_passed(params_partial_credentials) is True + + +def test_get_dynamic_langfuse_logging_config(): + dynamic_params = StandardCallbackDynamicParams( + langfuse_public_key="dynamic_key", + langfuse_secret="dynamic_secret", + langfuse_host="https://dynamic.langfuse.com", + ) + config = LangFuseHandler.get_dynamic_langfuse_logging_config(dynamic_params) + assert config["langfuse_public_key"] == "dynamic_key" + assert config["langfuse_secret"] == "dynamic_secret" + assert config["langfuse_host"] == "https://dynamic.langfuse.com" + + empty_params = StandardCallbackDynamicParams() + config = LangFuseHandler.get_dynamic_langfuse_logging_config(empty_params) + assert config["langfuse_public_key"] is None + assert config["langfuse_secret"] is None + assert config["langfuse_host"] is None + + +def test_return_global_langfuse_logger(): + mock_cache = Mock() + global_logger = LangFuseLogger(langfuse_public_key="global_key", langfuse_secret="global_secret") + + result = LangFuseHandler._return_global_langfuse_logger(global_logger, mock_cache) + assert result == global_logger + + mock_cache.get_cache.return_value = global_logger + result = LangFuseHandler._return_global_langfuse_logger(None, mock_cache) + assert result == global_logger + + mock_cache.get_cache.return_value = None + with patch.object( + LangFuseHandler, + "_create_langfuse_logger_from_credentials", + return_value=global_logger, + ): + result = LangFuseHandler._return_global_langfuse_logger(None, mock_cache) + assert result == global_logger + + +def test_get_langfuse_logger_for_request_with_cached_logger(): + mock_cache = Mock() + cached_logger = LangFuseLogger(langfuse_public_key="cached_key", langfuse_secret="cached_secret") + mock_cache.get_cache.return_value = cached_logger + + dynamic_params = StandardCallbackDynamicParams( + langfuse_public_key="test_key", + langfuse_secret="test_secret", + langfuse_host="https://test.langfuse.com", + ) + + result = LangFuseHandler.get_langfuse_logger_for_request( + standard_callback_dynamic_params=dynamic_params, + in_memory_dynamic_logger_cache=mock_cache, + globalLangfuseLogger=None, + ) + + assert result == cached_logger + mock_cache.get_cache.assert_called_once() diff --git a/tests/unit/integrations/test_langfuse.py b/tests/unit/integrations/test_langfuse.py index 3c0238372a5..7801fe44d45 100644 --- a/tests/unit/integrations/test_langfuse.py +++ b/tests/unit/integrations/test_langfuse.py @@ -8,12 +8,25 @@ import unittest from typing import Final, Optional from unittest.mock import MagicMock, patch +import httpx import pytest +import respx import litellm from litellm.integrations.langfuse import langfuse as langfuse_module from litellm.integrations.langfuse.langfuse import LangFuseLogger from litellm.integrations.langfuse.langfuse_sdk import resolve_trace_id +from litellm.types.utils import ( + Choices, + Message, + ModelResponse, + StandardLoggingHiddenParams, + StandardLoggingMetadata, + StandardLoggingModelInformation, + StandardLoggingPayload, + TextChoices, + TextCompletionResponse, +) # Import LangfuseUsageDetails directly from the module where it's defined @@ -2520,3 +2533,182 @@ def test_returned_generation_id_names_the_exported_observation(): span = _exported_span(logger, exporter) assert returned["generation_id"] == format(span.context.span_id, "016x") + + +def create_standard_logging_payload() -> StandardLoggingPayload: + return StandardLoggingPayload( + id="test_id", + call_type="completion", + response_cost=0.1, + response_cost_failure_debug_info=None, + status="success", + total_tokens=30, + prompt_tokens=20, + completion_tokens=10, + startTime=1234567890.0, + endTime=1234567891.0, + completionStartTime=1234567890.5, + model_map_information=StandardLoggingModelInformation(model_map_key="gpt-5-mini", model_map_value=None), + model="gpt-5-mini", + model_id="model-123", + model_group="openai-gpt", + api_base="https://api.openai.com", + metadata=StandardLoggingMetadata( + user_api_key_hash="test_hash", + user_api_key_org_id=None, + user_api_key_alias="test_alias", + user_api_key_team_id="test_team", + user_api_key_user_id="test_user", + user_api_key_team_alias="test_team_alias", + spend_logs_metadata=None, + requester_ip_address="127.0.0.1", + requester_metadata=None, + ), + cache_hit=False, + cache_key=None, + saved_cache_cost=0.0, + request_tags=[], + end_user=None, + requester_ip_address="127.0.0.1", + messages=[{"role": "user", "content": "Hello, world!"}], + response={"choices": [{"message": {"content": "Hi there!"}}]}, + error_str=None, + model_parameters={"stream": True}, + hidden_params=StandardLoggingHiddenParams( + model_id="model-123", + cache_key=None, + api_base="https://api.openai.com", + response_cost="0.1", + additional_headers=None, + ), + ) + + +@pytest.fixture +def global_langfuse_logger() -> LangFuseLogger: + with respx.mock(assert_all_called=False) as router: + router.route().mock(side_effect=httpx.ConnectError("langfuse is unreachable")) + return LangFuseLogger( + langfuse_public_key="global_public_key", + langfuse_secret="global_secret", + langfuse_host="https://global.langfuse.com", + ) + + +def test_get_langfuse_tags(global_langfuse_logger): + mock_payload = create_standard_logging_payload() + mock_payload["request_tags"] = ["tag1", "tag2", "test_tag"] + + result = global_langfuse_logger._get_langfuse_tags(mock_payload) + assert result == ["tag1", "tag2", "test_tag"] + + mock_payload["request_tags"] = None + result = global_langfuse_logger._get_langfuse_tags(mock_payload) + assert result == [] + + mock_payload["request_tags"] = [] + result = global_langfuse_logger._get_langfuse_tags(mock_payload) + assert result == [] + + +def test_get_chat_content_for_langfuse(): + mock_response = ModelResponse(choices=[Choices(message=Message(role="assistant", content="Hello world"))]) + + result = LangFuseLogger._get_chat_content_for_langfuse(mock_response) + assert result["content"] == "Hello world" + assert result["role"] == "assistant" + + mock_response = ModelResponse(choices=[]) + result = LangFuseLogger._get_chat_content_for_langfuse(mock_response) + assert result is None + + +def test_get_text_completion_content_for_langfuse(): + mock_response = TextCompletionResponse(choices=[TextChoices(text="Hello world")]) + result = LangFuseLogger._get_text_completion_content_for_langfuse(mock_response) + assert result == "Hello world" + + mock_response = TextCompletionResponse(choices=[]) + result = LangFuseLogger._get_text_completion_content_for_langfuse(mock_response) + assert result is None + + mock_response = TextCompletionResponse() + result = LangFuseLogger._get_text_completion_content_for_langfuse(mock_response) + assert result is None + + +def test_apply_masking_function_with_string(): + import re + + def mask_credit_cards(data): + if isinstance(data, str): + return re.sub(r"\b\d{4}[\s-]?\d{4}[\s-]?\d{4}[\s-]?\d{4}\b", "[CARD]", data) + return data + + input_str = "My card is 4532-1234-5678-9012" + result = LangFuseLogger._apply_masking_function(input_str, mask_credit_cards) + assert result == "My card is [CARD]" + assert "4532" not in result + + input_str = "Hello world" + result = LangFuseLogger._apply_masking_function(input_str, mask_credit_cards) + assert result == "Hello world" + + +def test_apply_masking_function_with_dict(): + import re + + def mask_emails(data): + if isinstance(data, str): + return re.sub(r"[\w\.-]+@[\w\.-]+", "[EMAIL]", data) + return data + + input_dict = {"messages": [{"role": "user", "content": "My email is test@example.com"}]} + result = LangFuseLogger._apply_masking_function(input_dict, mask_emails) + assert result["messages"][0]["content"] == "My email is [EMAIL]" + assert "test@example.com" not in str(result) + + +def test_apply_masking_function_with_none(): + def dummy_mask(data): + return data + + result = LangFuseLogger._apply_masking_function(None, dummy_mask) + assert result is None + + +def test_apply_masking_function_with_list(): + import re + + def mask_ssn(data): + if isinstance(data, str): + return re.sub(r"\b\d{3}-\d{2}-\d{4}\b", "[SSN]", data) + return data + + input_list = ["SSN: 123-45-6789", "No sensitive data here"] + result = LangFuseLogger._apply_masking_function(input_list, mask_ssn) + assert result[0] == "SSN: [SSN]" + assert result[1] == "No sensitive data here" + + +def test_langfuse_v2_uses_standard_logging_model_parameters(): + standard_logging_object = create_standard_logging_payload() + standard_logging_object["model_parameters"] = {"temperature": 0.5, "stream": True} + + optional_params_with_secrets = { + "temperature": 0.5, + "api_key": "sk-secret-key-12345", + "secret_fields": {"raw_headers": {"Authorization": "Bearer sk-secret"}}, + } + + sanitized = standard_logging_object.get("model_parameters", optional_params_with_secrets) + assert "api_key" not in sanitized + assert "secret_fields" not in sanitized + assert sanitized["temperature"] == 0.5 + + from litellm.litellm_core_utils.model_param_helper import ModelParamHelper + + fallback_sanitized = ModelParamHelper.get_standard_logging_model_parameters(optional_params_with_secrets) + assert "api_key" not in fallback_sanitized + assert "secret_fields" not in fallback_sanitized + assert fallback_sanitized["temperature"] == 0.5 diff --git a/tests/unit/integrations/test_sqs.py b/tests/unit/integrations/test_sqs.py new file mode 100644 index 00000000000..5fd923362af --- /dev/null +++ b/tests/unit/integrations/test_sqs.py @@ -0,0 +1,237 @@ +import asyncio +import base64 +import json +import os +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.integrations.sqs import SQSLogger + + +@pytest.mark.asyncio +async def test_async_send_batch_does_not_await_send_directly(monkeypatch): + monkeypatch.setattr("litellm.aws_sqs_callback_params", {}) + spawned = [] + real_create_task = asyncio.create_task + + def spy_create_task(coro, *args, **kwargs): + task = real_create_task(coro, *args, **kwargs) + spawned.append(task) + return task + + monkeypatch.setattr(asyncio, "create_task", spy_create_task) + + logger = SQSLogger(sqs_queue_url="https://example.com", sqs_region_name="us-west-2") + logger.async_send_message = AsyncMock() + logger.log_queue = [{"log": 1}, {"log": 2}] + + try: + await logger.async_send_batch() + assert logger.async_send_message.await_count == 0 + finally: + for task in spawned: + task.cancel() + await asyncio.gather(*spawned, return_exceptions=True) + + +def test_sqs_logger_init_without_encryption(monkeypatch): + monkeypatch.setattr("litellm.aws_sqs_callback_params", {}) + monkeypatch.setattr(asyncio, "create_task", MagicMock()) + logger = SQSLogger(sqs_queue_url="https://example.com", sqs_region_name="us-west-2") + assert logger.sqs_queue_url == "https://example.com" + assert logger.app_crypto is None + + +def test_sqs_logger_init_with_encryption(monkeypatch): + monkeypatch.setattr("litellm.aws_sqs_callback_params", {}) + monkeypatch.setattr(asyncio, "create_task", MagicMock()) + key_b64 = base64.b64encode(os.urandom(32)).decode() + + logger = SQSLogger( + sqs_queue_url="https://example.com", + sqs_region_name="us-west-2", + sqs_aws_use_application_level_encryption=True, + sqs_app_encryption_key_b64=key_b64, + sqs_app_encryption_aad="tenant=bill", + ) + assert logger.app_crypto is not None + assert logger.sqs_app_encryption_aad == "tenant=bill" + + +def test_sqs_logger_init_with_encryption_missing_key(monkeypatch): + monkeypatch.setattr("litellm.aws_sqs_callback_params", {}) + monkeypatch.setattr(asyncio, "create_task", MagicMock()) + with pytest.raises(ValueError, match="required when encryption is enabled"): + SQSLogger( + sqs_queue_url="https://example.com", + sqs_region_name="us-west-2", + sqs_aws_use_application_level_encryption=True, + ) + + +@pytest.mark.asyncio +async def test_async_log_success_event_adds_to_queue(monkeypatch): + monkeypatch.setattr("litellm.aws_sqs_callback_params", {}) + monkeypatch.setattr(asyncio, "create_task", MagicMock()) + logger = SQSLogger(sqs_queue_url="https://example.com", sqs_region_name="us-west-2") + + fake_payload = {"some": "data"} + await logger.async_log_success_event({"standard_logging_object": fake_payload}, None, None, None) + assert fake_payload in logger.log_queue + + +@pytest.mark.asyncio +async def test_async_log_failure_event_adds_to_queue(monkeypatch): + monkeypatch.setattr("litellm.aws_sqs_callback_params", {}) + monkeypatch.setattr(asyncio, "create_task", MagicMock()) + logger = SQSLogger(sqs_queue_url="https://example.com", sqs_region_name="us-west-2") + + fake_payload = {"fail": True} + await logger.async_log_failure_event({"standard_logging_object": fake_payload}, None, None, None) + assert fake_payload in logger.log_queue + + +@pytest.mark.asyncio +async def test_strip_base64_removes_file_and_nontext_entries(): + logger = SQSLogger(sqs_strip_base64_files=True) + + payload = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "Hello world"}, + { + "type": "image", + "file": {"file_data": "data:image/png;base64,AAAA"}, + }, + { + "type": "file", + "file": {"file_data": "data:application/pdf;base64,BBBB"}, + }, + ], + }, + { + "role": "assistant", + "content": [ + {"type": "text", "text": "Response"}, + { + "type": "audio", + "file": {"file_data": "data:audio/wav;base64,CCCC"}, + }, + ], + }, + ] + } + + stripped = await logger._strip_base64_from_messages(payload) + + assert len(stripped["messages"][0]["content"]) == 1 + assert stripped["messages"][0]["content"][0]["text"] == "Hello world" + + assert len(stripped["messages"][1]["content"]) == 1 + assert stripped["messages"][1]["content"][0]["text"] == "Response" + + for msg in stripped["messages"]: + for content in msg["content"]: + assert "file" not in content + assert content.get("type") == "text" + + +@pytest.mark.asyncio +async def test_strip_base64_handles_empty_or_missing_messages(): + logger = SQSLogger(sqs_strip_base64_files=True) + + payload_no_messages = {} + stripped1 = await logger._strip_base64_from_messages(payload_no_messages) + assert stripped1 == payload_no_messages + + payload_empty = {"messages": []} + stripped2 = await logger._strip_base64_from_messages(payload_empty) + assert stripped2 == payload_empty + + +@pytest.mark.asyncio +async def test_strip_base64_mixed_nested_objects(): + logger = SQSLogger(sqs_strip_base64_files=True) + + payload = { + "messages": [ + { + "role": "system", + "content": [ + {"type": "text", "text": "Keep me"}, + {"type": "custom", "metadata": "ignore but non-text"}, + {"foo": "bar"}, + {"file": {"file_data": "data:application/pdf;base64,XXX"}}, + ], + "extra": {"trace_id": "123"}, + } + ] + } + + stripped = await logger._strip_base64_from_messages(payload) + + content = stripped["messages"][0]["content"] + assert len(content) == 2 + assert {"type": "text", "text": "Keep me"} in content + assert {"foo": "bar"} in content + assert stripped["messages"][0]["extra"]["trace_id"] == "123" + + +@pytest.mark.asyncio +async def test_strip_base64_recursive_redaction(): + logger = SQSLogger(sqs_strip_base64_files=True) + payload = { + "messages": [ + { + "content": [ + {"type": "text", "text": "normal text"}, + { + "type": "text", + "text": "data:image/png;base64,iVBORw0KGgoAAAANSUhEUg", + }, + { + "type": "text", + "text": "Nested: {'data': 'data:application/pdf;base64,AAA...'}", + }, + {"file": {"file_data": "data:application/pdf;base64,AAAA"}}, + {"metadata": {"preview": "data:audio/mp3;base64,AAAAA=="}}, + ] + } + ] + } + + result = await logger._strip_base64_from_messages(payload) + content = result["messages"][0]["content"] + + assert not any("file" in c for c in content) + for c in content: + if isinstance(c, dict): + s = json.dumps(c).lower() + assert "base64," not in s, f"Found real base64 blob in: {s}" + + +@pytest.mark.asyncio +async def test_async_health_check_healthy(monkeypatch): + monkeypatch.setattr("litellm.aws_sqs_callback_params", {}) + monkeypatch.setattr(asyncio, "create_task", MagicMock()) + logger = SQSLogger(sqs_queue_url="https://example.com", sqs_region_name="us-west-2") + logger.async_send_message = AsyncMock(return_value=None) + + result = await logger.async_health_check() + assert result["status"] == "healthy" + assert result.get("error_message") is None + + +@pytest.mark.asyncio +async def test_async_health_check_unhealthy(monkeypatch): + monkeypatch.setattr("litellm.aws_sqs_callback_params", {}) + monkeypatch.setattr(asyncio, "create_task", MagicMock()) + logger = SQSLogger(sqs_queue_url="https://example.com", sqs_region_name="us-west-2") + logger.async_send_message = AsyncMock(side_effect=Exception("boom")) + + result = await logger.async_health_check() + assert result["status"] == "unhealthy" + assert "boom" in (result.get("error_message") or "") diff --git a/tests/unit/litellm_core_utils/llm_response_utils/test_convert_dict_to_response.py b/tests/unit/litellm_core_utils/llm_response_utils/test_convert_dict_to_response.py index 88e5e35b748..44d823fa82f 100644 --- a/tests/unit/litellm_core_utils/llm_response_utils/test_convert_dict_to_response.py +++ b/tests/unit/litellm_core_utils/llm_response_utils/test_convert_dict_to_response.py @@ -2814,3 +2814,78 @@ class TestConvertToModelResponseObjectCompletion: }, model_response_object=None, ) + + +def test_convert_to_model_response_object_function_output(): + response_object = { + "id": "chatcmpl-abc123", + "object": "chat.completion", + "created": 1699896916, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_abc123", + "type": "function", + "function": { + "name": "get_current_weather", + "arguments": '{\n"location": "Boston, MA"\n}', + }, + } + ], + }, + "logprobs": None, + "finish_reason": "tool_calls", + } + ], + "usage": { + "prompt_tokens": 82, + "completion_tokens": 17, + "total_tokens": 99, + "completion_tokens_details": {"reasoning_tokens": 0}, + }, + } + + result = convert_to_model_response_object( + model_response_object=ModelResponse(), + response_object=response_object, + stream=False, + start_time=datetime.now(), + end_time=datetime.now(), + hidden_params=None, + _response_headers=None, + convert_tool_call_to_json_mode=False, + ) + + assert isinstance(result, ModelResponse) + assert result.id == "chatcmpl-abc123" + assert result.object == "chat.completion" + assert result.created == 1699896916 + assert result.model == "gpt-4o-mini" + + assert len(result.choices) == 1 + choice = result.choices[0] + assert choice.index == 0 + assert isinstance(choice.message, Message) + assert choice.message.role == "assistant" + assert choice.message.content is None + assert choice.finish_reason == "tool_calls" + + assert len(choice.message.tool_calls) == 1 + tool_call = choice.message.tool_calls[0] + assert tool_call.id == "call_abc123" + assert tool_call.type == "function" + assert tool_call.function.name == "get_current_weather" + assert tool_call.function.arguments == '{\n"location": "Boston, MA"\n}' + + assert result.usage.prompt_tokens == 82 + assert result.usage.completion_tokens == 17 + assert result.usage.total_tokens == 99 + assert result.usage.completion_tokens_details == CompletionTokensDetailsWrapper(reasoning_tokens=0) + + assert result._hidden_params is not None diff --git a/tests/unit/litellm_core_utils/test_app_crypto.py b/tests/unit/litellm_core_utils/test_app_crypto.py new file mode 100644 index 00000000000..e0d68ef0568 --- /dev/null +++ b/tests/unit/litellm_core_utils/test_app_crypto.py @@ -0,0 +1,20 @@ +import os + +import pytest + +from litellm.litellm_core_utils.app_crypto import AppCrypto + + +def test_appcrypto_encrypt_decrypt_roundtrip(): + key = os.urandom(32) + crypto = AppCrypto(key) + data = {"event": "test", "value": 42} + aad = b"context" + enc = crypto.encrypt_json(data, aad=aad) + dec = crypto.decrypt_json(enc, aad=aad) + assert dec == data + + +def test_appcrypto_invalid_key_length(): + with pytest.raises(ValueError, match="32 bytes"): + AppCrypto(b"short") diff --git a/tests/unit/litellm_core_utils/test_litellm_logging.py b/tests/unit/litellm_core_utils/test_litellm_logging.py index caa8708ee9a..771ed5d2e15 100644 --- a/tests/unit/litellm_core_utils/test_litellm_logging.py +++ b/tests/unit/litellm_core_utils/test_litellm_logging.py @@ -11563,3 +11563,45 @@ class CompletionCustomHandler( except Exception: print(f"Assertion Error: {traceback.format_exc()}") self.errors.append(traceback.format_exc()) + + +def test_masking_function_isolated_from_other_loggers(): + from litellm.litellm_core_utils.litellm_logging import ( + scrub_sensitive_keys_in_metadata, + ) + + def my_masking_fn(data): + return data + + litellm_params = { + "metadata": { + "langfuse_masking_function": my_masking_fn, + "other_key": "other_value", + } + } + + result = scrub_sensitive_keys_in_metadata(litellm_params) + + assert "langfuse_masking_function" not in result["metadata"] + + assert result.get("_langfuse_masking_function") == my_masking_fn + + assert result["metadata"]["other_key"] == "other_value" + + +def test_masking_function_not_in_metadata_when_not_provided(): + from litellm.litellm_core_utils.litellm_logging import ( + scrub_sensitive_keys_in_metadata, + ) + + litellm_params = { + "metadata": { + "some_key": "some_value", + } + } + + result = scrub_sensitive_keys_in_metadata(litellm_params) + + assert "_langfuse_masking_function" not in result + + assert result["metadata"]["some_key"] == "some_value" diff --git a/tests/unit/litellm_core_utils/test_logging_callback_manager.py b/tests/unit/litellm_core_utils/test_logging_callback_manager.py index 31f65c3555e..f8260702068 100644 --- a/tests/unit/litellm_core_utils/test_logging_callback_manager.py +++ b/tests/unit/litellm_core_utils/test_logging_callback_manager.py @@ -5,6 +5,8 @@ import pytest import litellm from litellm.integrations.custom_logger import CustomLogger +from litellm.integrations.langfuse.langfuse_prompt_management import LangfusePromptManagement +from litellm.integrations.opentelemetry import OpenTelemetry from litellm.litellm_core_utils.logging_callback_manager import LoggingCallbackManager @@ -316,3 +318,27 @@ async def test_generic_api_callback_settings_retry_config(): finally: litellm.callback_settings.pop(callback_name, None) _generic_api_logger_cache.pop(callback_name, None) + + +def test_duplicate_langfuse_logger_test(): + manager = LoggingCallbackManager() + for _ in range(10): + langfuse_logger = LangfusePromptManagement() + manager.add_litellm_success_callback(langfuse_logger) + assert len(litellm.success_callback) == 1 + + +def test_duplicate_multiple_loggers_test(): + manager = LoggingCallbackManager() + for _ in range(10): + langfuse_logger = LangfusePromptManagement() + otel_logger = OpenTelemetry() + manager.add_litellm_success_callback(langfuse_logger) + manager.add_litellm_success_callback(otel_logger) + assert len(litellm.success_callback) == 2 + + langfuse_count = sum(1 for callback in litellm.success_callback if isinstance(callback, LangfusePromptManagement)) + otel_count = sum(1 for callback in litellm.success_callback if isinstance(callback, OpenTelemetry)) + + assert langfuse_count == 1, "Should have exactly one LangfusePromptManagement instance" + assert otel_count == 1, "Should have exactly one OpenTelemetry instance" diff --git a/tests/unit/litellm_core_utils/test_model_param_helper.py b/tests/unit/litellm_core_utils/test_model_param_helper.py index 2da324158cf..6293c748505 100644 --- a/tests/unit/litellm_core_utils/test_model_param_helper.py +++ b/tests/unit/litellm_core_utils/test_model_param_helper.py @@ -19,3 +19,30 @@ def test_get_all_llm_api_params_is_memoized(): first = ModelParamHelper.get_all_llm_api_params() second = ModelParamHelper.get_all_llm_api_params() assert first is second + + +def test_langfuse_model_parameters_no_secret_leakage(): + optional_params_with_secrets = { + "temperature": 0.7, + "top_p": 0.9, + "max_tokens": 100, + "stream": True, + "api_key": "sk-secret-key-12345", + "api_base": "https://my-private-endpoint.com", + "secret_fields": {"raw_headers": {"Authorization": "Bearer sk-super-secret"}}, + "authorization": "Bearer sk-another-secret", + "headers": {"X-Api-Key": "secret-header-value"}, + } + + sanitized = ModelParamHelper.get_standard_logging_model_parameters(optional_params_with_secrets) + + assert sanitized["temperature"] == 0.7 + assert sanitized["top_p"] == 0.9 + assert sanitized["max_tokens"] == 100 + assert sanitized["stream"] is True + + assert "api_key" not in sanitized + assert "api_base" not in sanitized + assert "secret_fields" not in sanitized + assert "authorization" not in sanitized + assert "headers" not in sanitized diff --git a/tests/unit/llms/bedrock/chat/invoke_transformations/test_base_invoke_transformation.py b/tests/unit/llms/bedrock/chat/invoke_transformations/test_base_invoke_transformation.py index d1748e1b38d..0a3cd9dc077 100644 --- a/tests/unit/llms/bedrock/chat/invoke_transformations/test_base_invoke_transformation.py +++ b/tests/unit/llms/bedrock/chat/invoke_transformations/test_base_invoke_transformation.py @@ -398,3 +398,22 @@ def test_router_deployment_with_a_non_numeric_stream_chunk_size_gets_a_400_befor assert exc_info.value.status_code == 400 send.assert_not_called() + + +@pytest.mark.respx(assert_all_called=False) +def test_transform_request_meta_llama(monkeypatch, respx_mock): + monkeypatch.setattr(litellm, "known_tokenizer_config", dict(litellm.known_tokenizer_config)) + respx_mock.get(host="huggingface.co").respond(404, text="Entry not found") + bedrock_transformer = AmazonInvokeConfig() + messages = [{"role": "user", "content": "Hello"}] + + result = bedrock_transformer.transform_request( + model="meta.llama2-70b", + messages=messages, + optional_params={"max_gen_len": 2048}, + litellm_params={}, + headers={}, + ) + + expected_result = {"prompt": "Hello", "max_gen_len": 2048} + assert result == expected_result diff --git a/tests/unit/llms/bedrock/embed/test_amazon_nova_transformation.py b/tests/unit/llms/bedrock/embed/test_amazon_nova_transformation.py index 6ca5f906f95..fa22c9b9bbc 100644 --- a/tests/unit/llms/bedrock/embed/test_amazon_nova_transformation.py +++ b/tests/unit/llms/bedrock/embed/test_amazon_nova_transformation.py @@ -537,3 +537,20 @@ def test_video_embedding_response_separate_mode(): assert len(result.data) == 2 assert result.data[0].embedding == [0.1, 0.2, 0.3] assert result.data[1].embedding == [0.4, 0.5, 0.6] + + +def test_async_invoke_requires_output_s3_uri(): + config = AmazonNovaEmbeddingConfig() + + inference_params = { + "embedding_purpose": "GENERIC_INDEX", + } + + with pytest.raises(ValueError, match="output_s3_uri is required"): + config.transform_request( + input="Test text", + inference_params=inference_params, + async_invoke_route=True, + model_id="amazon.nova-2-multimodal-embeddings-v1:0", + output_s3_uri=None, + ) diff --git a/tests/unit/llms/cloudflare/test_cloudflare_transformation.py b/tests/unit/llms/cloudflare/test_cloudflare_transformation.py index 13454cb46b6..3a45c83556c 100644 --- a/tests/unit/llms/cloudflare/test_cloudflare_transformation.py +++ b/tests/unit/llms/cloudflare/test_cloudflare_transformation.py @@ -387,3 +387,67 @@ def _make_mock_response(json_data: Dict[str, Any]) -> MagicMock: mock.json.return_value = json_data mock.text = json.dumps(json_data) return mock + + +@pytest.mark.parametrize("sync_mode", [False]) +def test_acompletion_cloudflare_stream(sync_mode, monkeypatch): + monkeypatch.setattr(litellm, "disable_hf_tokenizer_download", True) + messages = [{"role": "user", "content": "what llm are you"}] + raw_chunks = _streaming_chunks() + + if sync_mode: + + def _iter_lines(): + for chunk in raw_chunks: + yield f"data: {chunk}" + yield "data: [DONE]" + + mock_resp = MagicMock() + mock_resp.iter_lines.return_value = _iter_lines() + mock_resp.status_code = 200 + mock_resp.headers = {"content-type": "text/event-stream"} + + with patch.object(HTTPHandler, "post", return_value=mock_resp) as mock_post: + response = completion( + model="cloudflare/@cf/meta/llama-2-7b-chat-int8", + messages=messages, + max_tokens=15, + stream=True, + api_base=FAKE_API_BASE, + api_key=FAKE_API_KEY, + ) + chunks_received = list(response) + mock_post.assert_called_once() + else: + + async def _aiter_lines(): + for chunk in raw_chunks: + yield f"data: {chunk}" + yield "data: [DONE]" + + mock_resp = MagicMock() + mock_resp.aiter_lines.return_value = _aiter_lines() + mock_resp.status_code = 200 + mock_resp.headers = {"content-type": "text/event-stream"} + + async def _run(): + with patch.object(AsyncHTTPHandler, "post", new_callable=AsyncMock, return_value=mock_resp) as mock_post: + resp = await acompletion( + model="cloudflare/@cf/meta/llama-2-7b-chat-int8", + messages=messages, + max_tokens=15, + stream=True, + api_base=FAKE_API_BASE, + api_key=FAKE_API_KEY, + ) + received = [] + async for chunk in resp: + received.append(chunk) + mock_post.assert_called_once() + return received + + chunks_received = asyncio.run(_run()) + + assert len(chunks_received) > 0 + content = "".join(c.choices[0].delta.content for c in chunks_received if c.choices[0].delta.content) + assert "language" in content.lower() diff --git a/tests/unit/llms/litellm_proxy/chat/test_litellm_proxy_chat_transformation.py b/tests/unit/llms/litellm_proxy/chat/test_litellm_proxy_chat_transformation.py index 16e98cc29ec..86fb6fa15d8 100644 --- a/tests/unit/llms/litellm_proxy/chat/test_litellm_proxy_chat_transformation.py +++ b/tests/unit/llms/litellm_proxy/chat/test_litellm_proxy_chat_transformation.py @@ -9,6 +9,7 @@ from openai import AsyncOpenAI, OpenAI from openai.types import CreateEmbeddingResponse, Embedding from openai.types.create_embedding_response import Usage import pytest +import re import litellm from litellm import completion @@ -624,3 +625,18 @@ async def test_litellm_gateway_image_generation_direct(is_async): # Verify the response structure assert response is not None assert hasattr(response, "data") or isinstance(response, dict) + + +@pytest.mark.respx(assert_all_called=False) +def test_litellm_gateway_from_sdk_with_thinking_param(respx_mock): + respx_mock.route(host="0.0.0.0", port=4000).mock(side_effect=httpx.ConnectError("[Errno 61] Connection refused")) + with pytest.raises(Exception, match=re.escape("Connection error.")) as exc_info: + response = litellm.completion( + model="litellm_proxy/anthropic.claude-sonnet-4-5-20250929-v1:0", + messages=[{"role": "user", "content": "Hello world"}], + api_base="http://0.0.0.0:4000", + api_key="sk-PIp1h0RekR", + thinking={"type": "enabled", "max_budget": 100}, + ) + e = exc_info.value + assert "Connection error." in str(e) diff --git a/tests/unit/llms/nvidia_nim/test_nvidia_nim.py b/tests/unit/llms/nvidia_nim/test_nvidia_nim.py index 681343973d2..e5298c7ba39 100644 --- a/tests/unit/llms/nvidia_nim/test_nvidia_nim.py +++ b/tests/unit/llms/nvidia_nim/test_nvidia_nim.py @@ -197,3 +197,43 @@ async def test_nvidia_nim_rerank_ranking_endpoint(): # Model name in body should NOT have "ranking/" prefix assert request_data["model"] == "nvidia/llama-3.2-nv-rerankqa-1b-v2" + + +def test_completion_nvidia_nim(): + from openai import OpenAI + + litellm.set_verbose = True + model_name = "nvidia_nim/databricks/dbrx-instruct" + client = OpenAI( + api_key="fake-api-key", + ) + + with patch.object(client.chat.completions.with_raw_response, "create") as mock_client: + try: + completion( + model=model_name, + messages=[ + { + "role": "user", + "content": "What's the weather like in Boston today in Fahrenheit?", + } + ], + presence_penalty=0.5, + frequency_penalty=0.1, + client=client, + ) + except Exception as e: + print(e) + + mock_client.assert_called_once() + request_body = mock_client.call_args.kwargs + + assert request_body["messages"] == [ + { + "role": "user", + "content": "What's the weather like in Boston today in Fahrenheit?", + }, + ] + assert request_body["model"] == "databricks/dbrx-instruct" + assert request_body["frequency_penalty"] == 0.1 + assert request_body["presence_penalty"] == 0.5 diff --git a/tests/unit/llms/replicate/chat/test_transformation.py b/tests/unit/llms/replicate/chat/test_transformation.py index bbab4d50c4d..0270d86d485 100644 --- a/tests/unit/llms/replicate/chat/test_transformation.py +++ b/tests/unit/llms/replicate/chat/test_transformation.py @@ -7,6 +7,7 @@ import pytest from pydantic import ValidationError from litellm.llms.replicate.chat.handler import async_completion +from litellm.llms.replicate.chat.handler import completion as replicate_completion from litellm.llms.replicate.chat.transformation import ReplicateConfig from litellm.types.utils import ModelResponse @@ -183,3 +184,91 @@ def test_transform_response_string_output(): api_key="test-key", ) assert result.choices[0].message.content == "Hello from DeepSeek" + + +DEEPSEEK_V3_TOKENIZER_CONFIG = { + "add_bos_token": True, + "add_eos_token": False, + "bos_token": {"__type": "AddedToken", "content": "<|begin▁of▁sentence|>", "lstrip": False, "rstrip": False}, + "eos_token": {"__type": "AddedToken", "content": "<|end▁of▁sentence|>", "lstrip": False, "rstrip": False}, + "model_max_length": 131072, + "tokenizer_class": "LlamaTokenizerFast", + "chat_template": ( + "{{ bos_token }}{% for message in messages %}" + "{% if message['role'] == 'system' %}{{ message['content'] }}" + "{% elif message['role'] == 'user' %}{{ '<|User|>' + message['content'] }}" + "{% elif message['role'] == 'assistant' %}{{ '<|Assistant|>' + message['content'] + eos_token }}" + "{% endif %}{% endfor %}{% if add_generation_prompt %}{{ '<|Assistant|>' }}{% endif %}" + ), +} + + +@pytest.mark.respx(assert_all_called=False) +@patch("litellm.llms.replicate.chat.handler.get_httpx_client") +def test_sync_completion_handles_starting_status(mock_get_client, monkeypatch, respx_mock): + monkeypatch.setattr(litellm, "known_tokenizer_config", dict(litellm.known_tokenizer_config)) + respx_mock.get("https://huggingface.co/deepseek-ai/deepseek-v3/raw/main/tokenizer_config.json").respond( + 200, json=DEEPSEEK_V3_TOKENIZER_CONFIG + ) + mock_client = Mock() + mock_get_client.return_value = mock_client + + post_response = Mock() + post_response.json.return_value = { + "id": "test-prediction-id", + "urls": { + "get": "https://api.replicate.com/v1/predictions/test-id", + "cancel": "https://api.replicate.com/v1/predictions/test-id/cancel", + }, + } + mock_client.post.return_value = post_response + + get_response_starting = Mock() + get_response_starting.status_code = 200 + get_response_starting.json.return_value = { + "id": "test-prediction-id", + "status": "starting", + "output": None, + } + + get_response_succeeded = Mock() + get_response_succeeded.status_code = 200 + get_response_succeeded.json.return_value = { + "id": "test-prediction-id", + "status": "succeeded", + "output": ["Hello", " DeepSeek!"], + } + get_response_succeeded.text = json.dumps(get_response_succeeded.json.return_value) + get_response_succeeded.headers = {} + + mock_client.get.side_effect = [get_response_starting, get_response_succeeded] + + model_response = litellm.ModelResponse() + model_response.choices = [litellm.Choices()] + model_response.choices[0].message = litellm.Message(content="") + + mock_logging = Mock() + mock_logging.post_call = Mock() + + with patch("time.sleep"): + result = replicate_completion( + model="deepseek-ai/deepseek-v3", + messages=[{"role": "user", "content": "Hi"}], + api_base="https://api.replicate.com", + model_response=model_response, + print_verbose=print, + optional_params={}, + litellm_params={}, + logging_obj=mock_logging, + api_key="test-key", + encoding=None, + headers={}, + ) + + assert result is not None + assert result.choices[0].message.content == "Hello DeepSeek!" + + assert mock_client.get.call_count >= 1 + assert json.loads(mock_client.post.call_args.kwargs["data"])["input"]["prompt"] == ( + "<|begin▁of▁sentence|><|User|>Hi<|Assistant|>" + ) diff --git a/tests/unit/llms/triton/test_triton.py b/tests/unit/llms/triton/test_triton.py index 9ddd5b1db3b..7c0de603a3e 100644 --- a/tests/unit/llms/triton/test_triton.py +++ b/tests/unit/llms/triton/test_triton.py @@ -2,6 +2,7 @@ import json import traceback from unittest.mock import MagicMock, patch +import httpx import litellm import pytest @@ -191,3 +192,115 @@ def test_completion_triton_infer_api(): print("exception", e) traceback.print_exc() pytest.fail(f"Error occurred: {e}") + + +LLAMA_3_CHAT_TEMPLATE = ( + "{{ bos_token }}{% for message in messages %}" + "<|start_header_id|>{{ message['role'] }}<|end_header_id|>\n\n{{ message['content'] }}<|eot_id|>" + "{% endfor %}{% if add_generation_prompt %}<|start_header_id|>assistant<|end_header_id|>\n\n{% endif %}" +) + + +@pytest.mark.parametrize("stream", [True, False]) +@pytest.mark.respx(assert_all_called=False) +def test_completion_triton_generate_api(stream, monkeypatch, respx_mock): + monkeypatch.setattr(litellm, "known_tokenizer_config", dict(litellm.known_tokenizer_config)) + respx_mock.get(host="huggingface.co").respond(404, text="Entry not found") + try: + mock_response = MagicMock() + if stream: + + def mock_iter_lines(): + mock_output = "".join( + [ + 'data: {"model_name":"ensemble","model_version":"1","sequence_end":false,"sequence_id":0,"sequence_start":false,"text_output":"' + + t + + '"}\n\n' + for t in ["I", " am", " an", " AI", " assistant"] + ] + ) + for out in mock_output.split("\n"): + yield out + + mock_response.iter_lines = mock_iter_lines + else: + + def return_val(): + return { + "text_output": "I am an AI assistant", + } + + mock_response.json = return_val + mock_response.status_code = 200 + + with patch( + "litellm.llms.custom_httpx.http_handler.HTTPHandler.post", + return_value=mock_response, + ) as mock_post: + response = litellm.completion( + model="triton/llama-3-8b-instruct", + messages=[{"role": "user", "content": "who are u?"}], + max_tokens=10, + timeout=5, + api_base="http://localhost:8000/generate", + stream=stream, + ) + + mock_post.assert_called_once() + + call_kwargs = mock_post.call_args.kwargs + + if stream: + assert call_kwargs["url"] == "http://localhost:8000/generate_stream" + else: + assert call_kwargs["url"] == "http://localhost:8000/generate" + + request_data = json.loads(call_kwargs["data"]) + + assert request_data["text_input"] == "who are u?" + assert request_data["parameters"]["max_tokens"] == 10 + + if stream: + tokens = ["I", " am", " an", " AI", " assistant", None] + idx = 0 + for chunk in response: + assert chunk.choices[0].delta.content == tokens[idx] + idx += 1 + assert idx == len(tokens) + else: + assert response.choices[0].message.content == "I am an AI assistant" + + except Exception as e: + print("exception", e) + traceback.print_exc() + pytest.fail(f"Error occurred: {e}") + + +@pytest.mark.respx(assert_all_called=False) +def test_triton_generate_raw_request(monkeypatch, respx_mock): + monkeypatch.setattr(litellm, "known_tokenizer_config", dict(litellm.known_tokenizer_config)) + respx_mock.get("https://huggingface.co/llama-3-8b-instruct/raw/main/tokenizer_config.json").respond( + 404, text="Entry not found" + ) + respx_mock.get("https://huggingface.co/llama-3-8b-instruct/raw/main/chat_template.jinja").respond( + 200, text=LLAMA_3_CHAT_TEMPLATE + ) + from litellm.utils import return_raw_request + from litellm.types.utils import CallTypes + + try: + kwargs = { + "model": "triton/llama-3-8b-instruct", + "messages": [{"role": "user", "content": "who are u?"}], + "api_base": "http://localhost:8000/generate", + } + raw_request = return_raw_request(endpoint=CallTypes.completion, kwargs=kwargs) + assert raw_request is not None + assert "bad_words" not in json.dumps(raw_request["raw_request_body"]) + assert "stop_words" not in json.dumps(raw_request["raw_request_body"]) + assert raw_request["raw_request_body"]["text_input"] == ( + "<|start_header_id|>user<|end_header_id|>\n\nwho are u?<|eot_id|>" + "<|start_header_id|>assistant<|end_header_id|>\n\n" + ) + except Exception as e: + pytest.fail(f"Error occurred: {e}") diff --git a/tests/unit/secret_managers/test_cyberark_secret_manager.py b/tests/unit/secret_managers/test_cyberark_secret_manager.py index 334e6437f1d..a22a7a66294 100644 --- a/tests/unit/secret_managers/test_cyberark_secret_manager.py +++ b/tests/unit/secret_managers/test_cyberark_secret_manager.py @@ -2,13 +2,16 @@ import asyncio import json from pathlib import Path from typing import Final, TypedDict, cast +from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest import respx +import yaml import litellm import litellm.proxy.proxy_server +from litellm._uuid import uuid from litellm.secret_managers.cyberark_secret_manager import CyberArkSecretManager FIXTURE_PATH: Final = Path(__file__).resolve().parents[3] / "litellm-rust/crates/secrets-cyberark/tests/fixtures/parity.json" @@ -199,3 +202,281 @@ def test_missing_credentials_raise_value_error(monkeypatch: pytest.MonkeyPatch) monkeypatch.delenv(name, raising=False) with pytest.raises(ValueError, match="Missing CyberArk credentials"): CyberArkSecretManager() + + +@pytest.fixture +def cyberark_env(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("CYBERARK_API_KEY", "test-cyberark-api-key-909") + monkeypatch.setenv("CYBERARK_API_BASE", "http://0.0.0.0:8080") + monkeypatch.setenv("CYBERARK_ACCOUNT", "default") + monkeypatch.setenv("CYBERARK_USERNAME", "admin") + + +def create_mock_response(status_code: int, text: str = ""): + mock_response = MagicMock() + mock_response.status_code = status_code + mock_response.text = text + mock_response.raise_for_status = MagicMock() + + if status_code >= 400: + error = httpx.HTTPStatusError(message=f"HTTP {status_code}", request=MagicMock(), response=mock_response) + mock_response.raise_for_status.side_effect = error + + return mock_response + + +@pytest.mark.asyncio +async def test_cyberark_write_secret_rejects_yaml_injection(cyberark_env): + with patch("litellm.proxy.proxy_server.premium_user", True): + malicious_secret_name = "foo\n- !grant\n role: !!admin\n member: attacker" + + mock_sync_client = MagicMock() + mock_async_client = AsyncMock() + + with ( + patch( + "litellm.secret_managers.cyberark_secret_manager.get_httpx_client", + return_value=mock_sync_client, + ), + patch( + "litellm.secret_managers.cyberark_secret_manager.get_async_httpx_client", + return_value=mock_async_client, + ), + ): + cyberark_manager = CyberArkSecretManager() + + response = await cyberark_manager.async_write_secret( + secret_name=malicious_secret_name, + secret_value="sk-9876", + ) + + assert response["status"] == "error" + assert "Invalid secret_name" in response["message"] + mock_sync_client.client.post.assert_not_called() + mock_async_client.post.assert_not_called() + + +@pytest.mark.parametrize( + "secret_name", + [ + "foo: bar", + "foo # bar", + "plain-alias", + "team/user@example.com", + ], +) +@pytest.mark.asyncio +async def test_cyberark_ensure_variable_exists_escapes_yaml_metacharacters(cyberark_env, secret_name): + with patch("litellm.proxy.proxy_server.premium_user", True): + captured = {} + + async def _capture_post(url, headers=None, content=None): + captured["content"] = content + return create_mock_response(status_code=201, text="") + + mock_sync_client = MagicMock() + mock_sync_client.client.post.return_value = create_mock_response(status_code=200, text="mock-token") + mock_async_client = MagicMock() + mock_async_client.client.post.side_effect = _capture_post + + with patch( + "litellm.secret_managers.cyberark_secret_manager.get_httpx_client", + return_value=mock_sync_client, + ): + cyberark_manager = CyberArkSecretManager() + await cyberark_manager._ensure_variable_exists(secret_name, mock_async_client) + + policy_yaml = captured["content"] + parsed = yaml.compose(policy_yaml) + assert len(parsed.value) == 1 + node = parsed.value[0] + assert node.tag == "!variable" + assert node.value == secret_name + + +@pytest.mark.asyncio +async def test_cyberark_write_and_read_secret(cyberark_env): + with patch("litellm.proxy.proxy_server.premium_user", True): + secret_name = f"test-secret-{uuid.uuid4()}" + secret_value = f"test-value-{uuid.uuid4()}" + + mock_sync_client = MagicMock() + mock_sync_client.client.post.return_value = create_mock_response(status_code=200, text="mock-token") + mock_sync_client.client.get.return_value = create_mock_response(status_code=200, text=secret_value) + + mock_async_client = AsyncMock() + mock_async_client.post.return_value = create_mock_response(status_code=201, text="") + + with ( + patch( + "litellm.secret_managers.cyberark_secret_manager.get_httpx_client", + return_value=mock_sync_client, + ), + patch( + "litellm.secret_managers.cyberark_secret_manager.get_async_httpx_client", + return_value=mock_async_client, + ), + ): + cyberark_manager = CyberArkSecretManager() + + write_response = await cyberark_manager.async_write_secret( + secret_name=secret_name, + secret_value=secret_value, + ) + + assert write_response["status"] == "success" + + read_value = cyberark_manager.sync_read_secret(secret_name=secret_name) + + assert read_value is not None + assert read_value == secret_value + + +@pytest.mark.asyncio +async def test_cyberark_rotate_secret(cyberark_env): + with patch("litellm.proxy.proxy_server.premium_user", True): + secret_alias = f"test-rotation-key-{uuid.uuid4()}" + initial_key_value = f"sk-initial-{uuid.uuid4()}" + rotated_key_value = f"sk-rotated-{uuid.uuid4()}" + + current_value = {"value": initial_key_value} + + mock_sync_client = MagicMock() + mock_sync_client.client.post.return_value = create_mock_response(status_code=200, text="mock-token") + + def get_mock_sync_read_response(*args, **kwargs): + return create_mock_response(status_code=200, text=current_value["value"]) + + mock_sync_client.client.get.side_effect = get_mock_sync_read_response + + mock_async_client = AsyncMock() + + async def mock_async_post(*args, **kwargs): + content = kwargs.get("content", "") + if content: + current_value["value"] = content + return create_mock_response(status_code=201, text="") + + mock_async_client.post.side_effect = mock_async_post + + async def get_mock_async_read_response(*args, **kwargs): + return create_mock_response(status_code=200, text=current_value["value"]) + + mock_async_client.get.side_effect = get_mock_async_read_response + + with ( + patch( + "litellm.secret_managers.cyberark_secret_manager.get_httpx_client", + return_value=mock_sync_client, + ), + patch( + "litellm.secret_managers.cyberark_secret_manager.get_async_httpx_client", + return_value=mock_async_client, + ), + ): + cyberark_manager = CyberArkSecretManager() + + write_response = await cyberark_manager.async_write_secret( + secret_name=secret_alias, + secret_value=initial_key_value, + ) + assert write_response["status"] == "success" + + initial_read = cyberark_manager.sync_read_secret(secret_name=secret_alias) + assert initial_read == initial_key_value + + rotation_response = await cyberark_manager.async_rotate_secret( + current_secret_name=secret_alias, + new_secret_name=secret_alias, + new_secret_value=rotated_key_value, + ) + assert rotation_response["status"] == "success" + + cyberark_manager.cache.flush_cache() + + rotated_read = cyberark_manager.sync_read_secret(secret_name=secret_alias) + + assert rotated_read is not None + assert rotated_read == rotated_key_value + assert rotated_read != initial_key_value + + +@pytest.mark.asyncio +async def test_cyberark_rotate_secret_with_new_alias(cyberark_env): + with patch("litellm.proxy.proxy_server.premium_user", True): + base_alias = f"test-alias-change-{uuid.uuid4()}" + old_alias = f"{base_alias}-v1" + new_alias = f"{base_alias}-v2" + old_value = f"sk-old-{uuid.uuid4()}" + new_value = f"sk-new-{uuid.uuid4()}" + + secrets_store = {} + + mock_sync_client = MagicMock() + mock_sync_client.client.post.return_value = create_mock_response(status_code=200, text="mock-token") + + def get_mock_sync_read(*args, **kwargs): + url = args[0] if args else kwargs.get("url", "") + for secret_name, secret_val in secrets_store.items(): + if secret_name in url: + return create_mock_response(status_code=200, text=secret_val) + return create_mock_response(status_code=404, text="Not found") + + mock_sync_client.client.get.side_effect = get_mock_sync_read + + mock_async_client = AsyncMock() + + async def mock_async_post(*args, **kwargs): + url = args[0] if args else kwargs.get("url", "") + content = kwargs.get("content", "") + + if old_alias in url: + secrets_store[old_alias] = content + elif new_alias in url: + secrets_store[new_alias] = content + + return create_mock_response(status_code=201, text="") + + mock_async_client.post.side_effect = mock_async_post + + async def get_mock_async_read(*args, **kwargs): + url = args[0] if args else kwargs.get("url", "") + for secret_name, secret_val in secrets_store.items(): + if secret_name in url: + return create_mock_response(status_code=200, text=secret_val) + return create_mock_response(status_code=404, text="Not found") + + mock_async_client.get.side_effect = get_mock_async_read + + with ( + patch( + "litellm.secret_managers.cyberark_secret_manager.get_httpx_client", + return_value=mock_sync_client, + ), + patch( + "litellm.secret_managers.cyberark_secret_manager.get_async_httpx_client", + return_value=mock_async_client, + ), + ): + cyberark_manager = CyberArkSecretManager() + + write_response = await cyberark_manager.async_write_secret( + secret_name=old_alias, + secret_value=old_value, + ) + assert write_response["status"] == "success" + + rotation_response = await cyberark_manager.async_rotate_secret( + current_secret_name=old_alias, + new_secret_name=new_alias, + new_secret_value=new_value, + ) + assert rotation_response["status"] == "success" + + cyberark_manager.cache.flush_cache() + + new_read = cyberark_manager.sync_read_secret(secret_name=new_alias) + assert new_read == new_value + + old_read = cyberark_manager.sync_read_secret(secret_name=old_alias) + assert old_read == old_value diff --git a/tests/unit/test_utils_get_model_info.py b/tests/unit/test_utils_get_model_info.py index 713225af862..09c9079d89e 100644 --- a/tests/unit/test_utils_get_model_info.py +++ b/tests/unit/test_utils_get_model_info.py @@ -6,6 +6,7 @@ from typing import Final, Literal, cast import httpx import pytest +import respx import litellm from litellm import get_model_info @@ -501,3 +502,33 @@ def test_get_model_info_custom_provider(): get_model_info( model="my-custom-llm/my-fake-model" ) # 💥 "Exception: This model isn't mapped yet." in v1.56.10 + + +def test_get_model_info_huggingface_models(monkeypatch): + from litellm import Router + from litellm.types.router import ModelGroupInfo + + monkeypatch.setenv("HUGGINGFACE_API_KEY", "hf_abc123") + + with respx.mock(assert_all_called=False) as huggingface: + huggingface.get(host="huggingface.co").mock(side_effect=httpx.ConnectError("huggingface.co is unreachable")) + router = Router( + model_list=[ + { + "model_name": "meta-llama/Meta-Llama-3-8B-Instruct", + "litellm_params": { + "model": "huggingface/meta-llama/Meta-Llama-3-8B-Instruct", + "api_base": "https://router.huggingface.co/hf-inference/models/meta-llama/Meta-Llama-3-8B-Instruct", + "api_key": os.environ["HUGGINGFACE_API_KEY"], + }, + } + ] + ) + info = litellm.get_model_info("huggingface/meta-llama/Meta-Llama-3-8B-Instruct") + assert info is not None + + ModelGroupInfo( + model_group="meta-llama/Meta-Llama-3-8B-Instruct", + providers=["huggingface"], + **info, + ) diff --git a/tests/unit/test_vcr_classification.py b/tests/unit/test_vcr_classification.py index 568c1f64af2..0c219e9edf1 100644 --- a/tests/unit/test_vcr_classification.py +++ b/tests/unit/test_vcr_classification.py @@ -1,5 +1,6 @@ from __future__ import annotations +import socket from types import SimpleNamespace from typing import Optional @@ -545,3 +546,34 @@ def test_should_skip_live_probe_when_vcr_active(vcr_enabled): fake_cassette = SimpleNamespace(play_count=0, dirty=False) probe = install_live_call_probe(request, fake_cassette) assert probe is None + + +def test_live_call_probe_records_known_llm_hosts(vcr_enabled, monkeypatch): + def refuse_connection(address, *args, **kwargs): + raise ConnectionRefusedError(f"unit tests do not open sockets: {address}") + + monkeypatch.setattr(socket, "create_connection", refuse_connection) + finalizers = [] + + class _Node: + pass + + request = SimpleNamespace(node=_Node(), addfinalizer=lambda fn: finalizers.append(fn)) + probe = install_live_call_probe(request, None) + assert probe is not None + + try: + socket.create_connection(("api.openai.com", 443), timeout=0.001) + except Exception: + pass + try: + socket.create_connection(("127.0.0.1", 6379), timeout=0.001) + except Exception: + pass + + for fn in finalizers: + fn() + + hosts = getattr(request.node, "vcr_live_call_hosts", []) + assert "api.openai.com" in hosts + assert "127.0.0.1" not in hosts