mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
3eeca75ca0
commit
3250ccc802
35 changed files with 1505 additions and 2142 deletions
|
|
@ -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 |
|
||||
| --- | --- |
|
||||
|
|
|
|||
|
|
@ -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)")
|
||||
|
|
@ -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"
|
||||
|
|
@ -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!")
|
||||
|
|
@ -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()
|
||||
|
|
@ -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)
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
||||
|
|
@ -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)
|
||||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -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)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 "")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
163
tests/unit/integrations/langfuse/test_langfuse_handler.py
Normal file
163
tests/unit/integrations/langfuse/test_langfuse_handler.py
Normal 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()
|
||||
|
|
@ -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
|
||||
|
|
|
|||
237
tests/unit/integrations/test_sqs.py
Normal file
237
tests/unit/integrations/test_sqs.py
Normal 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 "")
|
||||
|
|
@ -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
|
||||
|
|
|
|||
20
tests/unit/litellm_core_utils/test_app_crypto.py
Normal file
20
tests/unit/litellm_core_utils/test_app_crypto.py
Normal 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")
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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|>"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue