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
This commit is contained in:
yuneng-jiang 2026-10-09 09:48:10 -07:00 • committed by GitHub
parent 3eeca75ca0
commit 3250ccc802
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
35 changed files with 1505 additions and 2142 deletions

View file

@ -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 |
| --- | --- |

View file

@ -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)")

View file

@ -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"

View file

@ -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!")

View file

@ -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()

View file

@ -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)

View file

@ -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

View file

@ -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

View file

@ -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)

View file

@ -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}")

View file

@ -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

View file

@ -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

View file

@ -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,
)

View file

@ -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)

View file

@ -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

View file

@ -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 "")

View file

@ -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)

View file

@ -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()

View file

@ -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

View file

@ -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 "")

View file

@ -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

View file

@ -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")

View file

@ -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"

View file

@ -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"

View file

@ -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

View file

@ -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

View file

@ -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,
)

View file

@ -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()

View file

@ -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)

View file

@ -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

View file

@ -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|>"
)

View file

@ -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}")

View file

@ -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

View file

@ -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,
)

View file

@ -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