mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
add tests for hotpath & docker container
This commit is contained in:
parent
f8660a8ab0
commit
8249cce157
2 changed files with 946 additions and 0 deletions
|
|
@ -0,0 +1,543 @@
|
|||
"""
|
||||
Test to count and track the number of network requests (DB queries, cache lookups)
|
||||
made on the hot path for keys that have team_id and user_id attached.
|
||||
|
||||
This test ensures we don't regress on the number of network requests made during
|
||||
request authentication, which directly impacts proxy latency.
|
||||
|
||||
The hot path covers auth functions called on every LLM API request:
|
||||
- get_key_object: lookup the API key
|
||||
- get_team_object: lookup the team (for keys with team_id)
|
||||
- get_user_object: lookup the user (for keys with user_id)
|
||||
- get_team_membership: lookup team member budget (when team_member_spend set)
|
||||
|
||||
Each function does: cache read -> (on miss) DB query -> cache write.
|
||||
We count these to catch regressions in the number of network requests.
|
||||
|
||||
NOTE: This test does NOT require proxy extras (apscheduler, etc.) because
|
||||
it tests at the auth_checks level, not the full proxy_server level.
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from typing import Any, Dict, List, Optional
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../.."))
|
||||
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_TeamTableCachedObj,
|
||||
LiteLLM_UserTable,
|
||||
LitellmUserRoles,
|
||||
UserAPIKeyAuth,
|
||||
LiteLLM_TeamMembership,
|
||||
hash_token,
|
||||
)
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
get_key_object,
|
||||
get_team_membership,
|
||||
get_team_object,
|
||||
get_user_object,
|
||||
)
|
||||
|
||||
class CacheCallTracker:
|
||||
"""
|
||||
Tracks cache read/write operations by wrapping DualCache methods.
|
||||
This is used to count network-level operations on the hot path.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.cache_reads: List[Dict[str, Any]] = []
|
||||
self.cache_writes: List[Dict[str, Any]] = []
|
||||
self.db_queries: List[Dict[str, Any]] = []
|
||||
|
||||
def get_summary(self) -> Dict[str, Any]:
|
||||
return {
|
||||
"total_cache_reads": len(self.cache_reads),
|
||||
"total_cache_writes": len(self.cache_writes),
|
||||
"total_db_queries": len(self.db_queries),
|
||||
"total_network_requests": len(self.cache_reads)
|
||||
+ len(self.cache_writes)
|
||||
+ len(self.db_queries),
|
||||
"cache_read_keys": [r["key"] for r in self.cache_reads],
|
||||
"cache_write_keys": [w["key"] for w in self.cache_writes],
|
||||
"db_query_details": self.db_queries,
|
||||
}
|
||||
|
||||
|
||||
def _wrap_cache_with_tracker(cache: DualCache, tracker: CacheCallTracker) -> DualCache:
|
||||
"""Wrap a DualCache to track all reads and writes."""
|
||||
original_async_get = cache.async_get_cache
|
||||
original_async_set = cache.async_set_cache
|
||||
|
||||
async def tracked_async_get(key, *args, **kwargs):
|
||||
result = await original_async_get(key, *args, **kwargs)
|
||||
tracker.cache_reads.append(
|
||||
{"key": key, "hit": result is not None, "method": "async_get_cache"}
|
||||
)
|
||||
return result
|
||||
|
||||
async def tracked_async_set(key, value, *args, **kwargs):
|
||||
tracker.cache_writes.append({"key": key, "method": "async_set_cache"})
|
||||
return await original_async_set(key, value, *args, **kwargs)
|
||||
|
||||
cache.async_get_cache = tracked_async_get
|
||||
cache.async_set_cache = tracked_async_set
|
||||
return cache
|
||||
|
||||
|
||||
def _create_valid_token(
|
||||
api_key: str,
|
||||
team_id: str,
|
||||
user_id: str,
|
||||
has_team_member_spend: bool = False,
|
||||
org_id: Optional[str] = None,
|
||||
) -> UserAPIKeyAuth:
|
||||
"""Create a UserAPIKeyAuth with team_id and user_id set."""
|
||||
hashed = hash_token(api_key)
|
||||
return UserAPIKeyAuth(
|
||||
token=hashed,
|
||||
api_key=api_key,
|
||||
team_id=team_id,
|
||||
user_id=user_id,
|
||||
org_id=org_id,
|
||||
models=["gpt-4", "gpt-3.5-turbo"],
|
||||
max_budget=100.0,
|
||||
spend=10.0,
|
||||
team_spend=50.0,
|
||||
team_max_budget=1000.0,
|
||||
team_models=["gpt-4", "gpt-3.5-turbo"],
|
||||
team_member_spend=5.0 if has_team_member_spend else None,
|
||||
last_refreshed_at=time.time(),
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
)
|
||||
|
||||
|
||||
def _create_team_object(team_id: str) -> LiteLLM_TeamTableCachedObj:
|
||||
"""Create a team table object for caching."""
|
||||
return LiteLLM_TeamTableCachedObj(
|
||||
team_id=team_id,
|
||||
models=["gpt-4", "gpt-3.5-turbo"],
|
||||
max_budget=1000.0,
|
||||
spend=50.0,
|
||||
tpm_limit=10000,
|
||||
rpm_limit=100,
|
||||
last_refreshed_at=time.time(),
|
||||
)
|
||||
|
||||
|
||||
def _create_user_object(user_id: str) -> LiteLLM_UserTable:
|
||||
"""Create a user table object for caching."""
|
||||
return LiteLLM_UserTable(
|
||||
user_id=user_id,
|
||||
max_budget=500.0,
|
||||
spend=25.0,
|
||||
models=["gpt-4"],
|
||||
tpm_limit=5000,
|
||||
rpm_limit=50,
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
user_email="test@example.com",
|
||||
)
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# TEST: get_key_object cache behavior
|
||||
# ============================================================================
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_key_object_warm_cache():
|
||||
"""
|
||||
Test get_key_object with a warm cache - should hit cache, no DB query.
|
||||
"""
|
||||
api_key = "sk-test-key-warm"
|
||||
team_id = "team-123"
|
||||
user_id = "user-456"
|
||||
hashed_token = hash_token(api_key)
|
||||
|
||||
valid_token = _create_valid_token(api_key, team_id, user_id)
|
||||
|
||||
# Create cache with pre-populated data
|
||||
cache = DualCache(in_memory_cache=InMemoryCache())
|
||||
await cache.async_set_cache(key=hashed_token, value=valid_token)
|
||||
|
||||
# Track cache operations
|
||||
tracker = CacheCallTracker()
|
||||
tracked_cache = _wrap_cache_with_tracker(cache, tracker)
|
||||
|
||||
# Mock prisma client (should NOT be called for warm cache)
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.get_data = AsyncMock()
|
||||
|
||||
result = await get_key_object(
|
||||
hashed_token=hashed_token,
|
||||
prisma_client=mock_prisma,
|
||||
user_api_key_cache=tracked_cache,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=None,
|
||||
)
|
||||
|
||||
summary = tracker.get_summary()
|
||||
|
||||
print("\n=== get_key_object WARM CACHE ===")
|
||||
print(f"Cache reads: {summary['total_cache_reads']}")
|
||||
print(f"Keys read: {summary['cache_read_keys']}")
|
||||
|
||||
# Should have exactly 1 cache read
|
||||
assert summary["total_cache_reads"] == 1
|
||||
assert hashed_token in summary["cache_read_keys"]
|
||||
|
||||
# Prisma should NOT have been called
|
||||
mock_prisma.get_data.assert_not_called()
|
||||
|
||||
# Result should be the cached token
|
||||
assert result.token == hashed_token
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_key_object_cold_cache():
|
||||
"""
|
||||
Test get_key_object with a cold cache - should miss cache, query DB.
|
||||
"""
|
||||
api_key = "sk-test-key-cold"
|
||||
team_id = "team-123"
|
||||
user_id = "user-456"
|
||||
hashed_token = hash_token(api_key)
|
||||
|
||||
valid_token = _create_valid_token(api_key, team_id, user_id)
|
||||
|
||||
# Create empty cache
|
||||
cache = DualCache(in_memory_cache=InMemoryCache())
|
||||
|
||||
tracker = CacheCallTracker()
|
||||
tracked_cache = _wrap_cache_with_tracker(cache, tracker)
|
||||
|
||||
# Mock prisma client to return token on DB query
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.get_data = AsyncMock(return_value=valid_token)
|
||||
|
||||
result = await get_key_object(
|
||||
hashed_token=hashed_token,
|
||||
prisma_client=mock_prisma,
|
||||
user_api_key_cache=tracked_cache,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=None,
|
||||
)
|
||||
|
||||
summary = tracker.get_summary()
|
||||
|
||||
print("\n=== get_key_object COLD CACHE ===")
|
||||
print(f"Cache reads: {summary['total_cache_reads']}")
|
||||
print(f"Cache writes: {summary['total_cache_writes']}")
|
||||
|
||||
# Should have 1 cache read (miss) and at least 1 cache write (populate cache)
|
||||
assert summary["total_cache_reads"] >= 1
|
||||
|
||||
# Prisma SHOULD have been called
|
||||
mock_prisma.get_data.assert_called_once()
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# TEST: get_team_object cache behavior
|
||||
# ============================================================================
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_team_object_warm_cache():
|
||||
"""
|
||||
Test get_team_object with a warm cache - should hit cache, no DB query.
|
||||
"""
|
||||
team_id = "team-warm-123"
|
||||
team_obj = _create_team_object(team_id)
|
||||
|
||||
cache = DualCache(in_memory_cache=InMemoryCache())
|
||||
cache_key = f"team_id:{team_id}"
|
||||
await cache.async_set_cache(key=cache_key, value=team_obj)
|
||||
|
||||
tracker = CacheCallTracker()
|
||||
tracked_cache = _wrap_cache_with_tracker(cache, tracker)
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db = MagicMock()
|
||||
mock_prisma.db.litellm_teamtable = MagicMock()
|
||||
mock_prisma.db.litellm_teamtable.find_unique = AsyncMock()
|
||||
|
||||
result = await get_team_object(
|
||||
team_id=team_id,
|
||||
prisma_client=mock_prisma,
|
||||
user_api_key_cache=tracked_cache,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=None,
|
||||
)
|
||||
|
||||
summary = tracker.get_summary()
|
||||
|
||||
print("\n=== get_team_object WARM CACHE ===")
|
||||
print(f"Cache reads: {summary['total_cache_reads']}")
|
||||
print(f"Keys read: {summary['cache_read_keys']}")
|
||||
|
||||
assert summary["total_cache_reads"] >= 1
|
||||
assert cache_key in summary["cache_read_keys"]
|
||||
|
||||
# DB should NOT have been called
|
||||
mock_prisma.db.litellm_teamtable.find_unique.assert_not_called()
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# TEST: get_user_object cache behavior
|
||||
# ============================================================================
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_user_object_warm_cache():
|
||||
"""
|
||||
Test get_user_object with a warm cache - should hit cache, no DB query.
|
||||
"""
|
||||
user_id = "user-warm-456"
|
||||
user_obj = _create_user_object(user_id)
|
||||
|
||||
cache = DualCache(in_memory_cache=InMemoryCache())
|
||||
await cache.async_set_cache(key=user_id, value=user_obj)
|
||||
|
||||
tracker = CacheCallTracker()
|
||||
tracked_cache = _wrap_cache_with_tracker(cache, tracker)
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db = MagicMock()
|
||||
mock_prisma.db.litellm_usertable = MagicMock()
|
||||
mock_prisma.db.litellm_usertable.find_unique = AsyncMock()
|
||||
|
||||
result = await get_user_object(
|
||||
user_id=user_id,
|
||||
prisma_client=mock_prisma,
|
||||
user_api_key_cache=tracked_cache,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=None,
|
||||
user_id_upsert=False,
|
||||
)
|
||||
|
||||
summary = tracker.get_summary()
|
||||
|
||||
print("\n=== get_user_object WARM CACHE ===")
|
||||
print(f"Cache reads: {summary['total_cache_reads']}")
|
||||
print(f"Keys read: {summary['cache_read_keys']}")
|
||||
|
||||
assert summary["total_cache_reads"] >= 1
|
||||
assert user_id in summary["cache_read_keys"]
|
||||
|
||||
# DB should NOT have been called
|
||||
mock_prisma.db.litellm_usertable.find_unique.assert_not_called()
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# TEST: get_team_membership cache behavior
|
||||
# ============================================================================
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_team_membership_warm_cache():
|
||||
"""
|
||||
Test get_team_membership with a warm cache - should hit cache, no DB query.
|
||||
"""
|
||||
user_id = "user-tm-456"
|
||||
team_id = "team-tm-123"
|
||||
|
||||
membership_dict = {
|
||||
"user_id": user_id,
|
||||
"team_id": team_id,
|
||||
"spend": 3.0,
|
||||
"budget_id": None,
|
||||
"litellm_budget_table": None,
|
||||
}
|
||||
|
||||
cache = DualCache(in_memory_cache=InMemoryCache())
|
||||
# Cache key format used by get_team_membership
|
||||
cache_key = f"team_membership:{user_id}:{team_id}"
|
||||
await cache.async_set_cache(key=cache_key, value=membership_dict)
|
||||
|
||||
tracker = CacheCallTracker()
|
||||
tracked_cache = _wrap_cache_with_tracker(cache, tracker)
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db = MagicMock()
|
||||
mock_prisma.db.litellm_teammembership = MagicMock()
|
||||
mock_prisma.db.litellm_teammembership.find_unique = AsyncMock()
|
||||
|
||||
result = await get_team_membership(
|
||||
user_id=user_id,
|
||||
team_id=team_id,
|
||||
prisma_client=mock_prisma,
|
||||
user_api_key_cache=tracked_cache,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=None,
|
||||
)
|
||||
|
||||
summary = tracker.get_summary()
|
||||
|
||||
print("\n=== get_team_membership WARM CACHE ===")
|
||||
print(f"Cache reads: {summary['total_cache_reads']}")
|
||||
print(f"Keys read: {summary['cache_read_keys']}")
|
||||
|
||||
assert summary["total_cache_reads"] >= 1
|
||||
assert cache_key in summary["cache_read_keys"]
|
||||
|
||||
# DB should NOT have been called
|
||||
mock_prisma.db.litellm_teammembership.find_unique.assert_not_called()
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# TEST: Document duplicate team membership cache key issue
|
||||
# ============================================================================
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_membership_cache_key_duplication():
|
||||
"""
|
||||
Document the team membership duplicate cache key issue:
|
||||
|
||||
Team membership is queried via TWO different cache keys:
|
||||
1. "{team_id}_{user_id}" - used in user_api_key_auth.py:1048
|
||||
2. "team_membership:{user_id}:{team_id}" - used in auth_checks.py:960 (get_team_membership)
|
||||
|
||||
This test documents that both keys refer to the same data but use different
|
||||
cache key formats, potentially leading to duplicate lookups.
|
||||
"""
|
||||
user_id = "user-dup-456"
|
||||
team_id = "team-dup-123"
|
||||
|
||||
# The two different cache keys used for the same data
|
||||
key_format_1 = f"{team_id}_{user_id}" # user_api_key_auth format
|
||||
key_format_2 = f"team_membership:{user_id}:{team_id}" # auth_checks format
|
||||
|
||||
membership_data = {
|
||||
"user_id": user_id,
|
||||
"team_id": team_id,
|
||||
"spend": 3.0,
|
||||
}
|
||||
|
||||
# Document that these are different keys
|
||||
assert key_format_1 != key_format_2, "Cache keys should be different (this is the bug)"
|
||||
|
||||
print("\n=== DOCUMENTATION: Team Membership Duplicate Cache Keys ===")
|
||||
print(f"Cache key format 1 (user_api_key_auth): {key_format_1}")
|
||||
print(f"Cache key format 2 (auth_checks): {key_format_2}")
|
||||
print("\nRecommendation: Consolidate to a single cache key format")
|
||||
print("to avoid duplicate cache reads/writes for the same data.")
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# TEST: Full hot path network count summary
|
||||
# ============================================================================
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_full_hot_path_network_count():
|
||||
"""
|
||||
Summary test that counts all network operations when processing
|
||||
a request with a key that has team_id and user_id attached.
|
||||
|
||||
This test verifies the baseline number of cache operations expected
|
||||
on a fully warm cache path.
|
||||
"""
|
||||
api_key = "sk-test-full-path"
|
||||
team_id = "team-full-123"
|
||||
user_id = "user-full-456"
|
||||
hashed_token = hash_token(api_key)
|
||||
|
||||
# Create all objects
|
||||
valid_token = _create_valid_token(api_key, team_id, user_id, has_team_member_spend=True)
|
||||
team_obj = _create_team_object(team_id)
|
||||
user_obj = _create_user_object(user_id)
|
||||
membership_data = LiteLLM_TeamMembership(
|
||||
user_id=user_id,
|
||||
team_id=team_id,
|
||||
spend=3.0,
|
||||
budget_id=None,
|
||||
litellm_budget_table=None,
|
||||
)
|
||||
|
||||
# Pre-populate cache with all data
|
||||
cache = DualCache(in_memory_cache=InMemoryCache())
|
||||
await cache.async_set_cache(key=hashed_token, value=valid_token)
|
||||
await cache.async_set_cache(key=f"team_id:{team_id}", value=team_obj)
|
||||
await cache.async_set_cache(key=user_id, value=user_obj)
|
||||
await cache.async_set_cache(key=f"team_membership:{user_id}:{team_id}", value=membership_data.model_dump())
|
||||
await cache.async_set_cache(key=f"{team_id}_{user_id}", value=membership_data.model_dump())
|
||||
|
||||
# Create tracker AFTER populating cache
|
||||
tracker = CacheCallTracker()
|
||||
tracked_cache = _wrap_cache_with_tracker(cache, tracker)
|
||||
|
||||
# Mock prisma (should not be called on warm cache)
|
||||
mock_prisma = MagicMock()
|
||||
|
||||
# Call each function to simulate the hot path
|
||||
await get_key_object(
|
||||
hashed_token=hashed_token,
|
||||
prisma_client=mock_prisma,
|
||||
user_api_key_cache=tracked_cache,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=None,
|
||||
)
|
||||
|
||||
await get_team_object(
|
||||
team_id=team_id,
|
||||
prisma_client=mock_prisma,
|
||||
user_api_key_cache=tracked_cache,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=None,
|
||||
)
|
||||
|
||||
await get_user_object(
|
||||
user_id=user_id,
|
||||
prisma_client=mock_prisma,
|
||||
user_api_key_cache=tracked_cache,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=None,
|
||||
user_id_upsert=False,
|
||||
)
|
||||
|
||||
await get_team_membership(
|
||||
user_id=user_id,
|
||||
team_id=team_id,
|
||||
prisma_client=mock_prisma,
|
||||
user_api_key_cache=tracked_cache,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=None,
|
||||
)
|
||||
|
||||
summary = tracker.get_summary()
|
||||
|
||||
print("\n" + "=" * 70)
|
||||
print("HOT PATH NETWORK REQUEST SUMMARY")
|
||||
print("Key with team_id and user_id attached - WARM CACHE")
|
||||
print("=" * 70)
|
||||
print(f"Cache reads: {summary['total_cache_reads']}")
|
||||
print(f"Cache writes: {summary['total_cache_writes']}")
|
||||
print(f"DB queries: {summary['total_db_queries']}")
|
||||
print(f"TOTAL: {summary['total_network_requests']}")
|
||||
print(f"\nCache keys read:")
|
||||
for key in summary["cache_read_keys"]:
|
||||
print(f" - {key}")
|
||||
print("=" * 70)
|
||||
|
||||
# Assertions for expected baseline
|
||||
# On warm cache: 4 reads (key, team, user, team_membership)
|
||||
assert summary["total_cache_reads"] == 4, (
|
||||
f"Expected 4 cache reads on warm path, got {summary['total_cache_reads']}"
|
||||
)
|
||||
|
||||
# No DB queries on warm cache
|
||||
assert summary["total_db_queries"] == 0, (
|
||||
f"Expected 0 DB queries on warm path, got {summary['total_db_queries']}"
|
||||
)
|
||||
|
||||
# Total network requests should be exactly 4 on warm cache
|
||||
assert summary["total_network_requests"] == 4, (
|
||||
f"Expected 4 total network requests on warm path, got {summary['total_network_requests']}"
|
||||
)
|
||||
403
tests/test_litellm/proxy/test_docker_no_network_on_deploy.py
Normal file
403
tests/test_litellm/proxy/test_docker_no_network_on_deploy.py
Normal file
|
|
@ -0,0 +1,403 @@
|
|||
"""
|
||||
Test to ensure Docker container does not go out to network on deploy.
|
||||
|
||||
This test verifies that the LiteLLM proxy container does not make outbound
|
||||
network requests during startup. This is important for:
|
||||
1. Air-gapped environments where outbound network is restricted
|
||||
2. Security compliance requiring no unexpected network calls
|
||||
3. Fast container startup without network dependencies
|
||||
|
||||
The test works by:
|
||||
1. Building/running the container with network disabled
|
||||
2. Verifying the container starts successfully without network
|
||||
3. Checking for any errors related to network failures during startup
|
||||
"""
|
||||
|
||||
import os
|
||||
import re
|
||||
import subprocess
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def is_docker_available() -> bool:
|
||||
"""Check if Docker is available and running."""
|
||||
try:
|
||||
result = subprocess.run(
|
||||
["docker", "info"],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=10,
|
||||
)
|
||||
return result.returncode == 0
|
||||
except (subprocess.TimeoutExpired, FileNotFoundError):
|
||||
return False
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not is_docker_available(),
|
||||
reason="Docker not available",
|
||||
)
|
||||
class TestDockerNoNetworkOnDeploy:
|
||||
"""
|
||||
Test suite for verifying Docker container starts without network access.
|
||||
"""
|
||||
|
||||
# Container and image names for testing
|
||||
TEST_IMAGE_NAME = "litellm-no-network-test"
|
||||
TEST_CONTAINER_NAME = "litellm-no-network-test-container"
|
||||
|
||||
# Timeout for container operations
|
||||
CONTAINER_START_TIMEOUT = 60 # seconds
|
||||
|
||||
# Patterns that indicate network-related failures during startup
|
||||
NETWORK_ERROR_PATTERNS = [
|
||||
r"connection refused",
|
||||
r"network is unreachable",
|
||||
r"could not resolve host",
|
||||
r"name resolution failed",
|
||||
r"dns lookup failed",
|
||||
r"failed to establish.*connection",
|
||||
r"socket.gaierror",
|
||||
r"urllib.error.URLError.*Errno",
|
||||
r"requests.exceptions.ConnectionError",
|
||||
r"httpx.*ConnectError",
|
||||
r"aiohttp.*ClientConnectorError",
|
||||
]
|
||||
|
||||
# Patterns that indicate EXPECTED local-only startup
|
||||
LOCAL_STARTUP_PATTERNS = [
|
||||
r"starting.*(server|proxy)",
|
||||
r"listening on",
|
||||
r"uvicorn running",
|
||||
r"application startup complete",
|
||||
]
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def cleanup(self):
|
||||
"""Cleanup any existing test containers before and after each test."""
|
||||
self._cleanup_container()
|
||||
yield
|
||||
self._cleanup_container()
|
||||
|
||||
def _cleanup_container(self):
|
||||
"""Remove test container if it exists."""
|
||||
subprocess.run(
|
||||
["docker", "rm", "-f", self.TEST_CONTAINER_NAME],
|
||||
capture_output=True,
|
||||
timeout=30,
|
||||
)
|
||||
|
||||
def _build_test_image(self) -> bool:
|
||||
"""
|
||||
Build the Docker image if needed.
|
||||
Returns True if image is available (built or already exists).
|
||||
"""
|
||||
# Check if main litellm image exists
|
||||
result = subprocess.run(
|
||||
["docker", "images", "-q", "litellm/litellm"],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=10,
|
||||
)
|
||||
if result.stdout.strip():
|
||||
# Image exists, use it
|
||||
return True
|
||||
|
||||
# Try to build from Dockerfile
|
||||
dockerfile_path = os.path.join(
|
||||
os.path.dirname(__file__),
|
||||
"..",
|
||||
"..",
|
||||
"..",
|
||||
"Dockerfile",
|
||||
)
|
||||
if os.path.exists(dockerfile_path):
|
||||
result = subprocess.run(
|
||||
[
|
||||
"docker",
|
||||
"build",
|
||||
"-t",
|
||||
self.TEST_IMAGE_NAME,
|
||||
"-f",
|
||||
dockerfile_path,
|
||||
os.path.dirname(dockerfile_path),
|
||||
],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=600, # 10 minutes for build
|
||||
)
|
||||
return result.returncode == 0
|
||||
|
||||
return False
|
||||
|
||||
def test_container_starts_without_network(self):
|
||||
"""
|
||||
Test that the container can start with network completely disabled.
|
||||
|
||||
This test runs the container with --network=none to ensure no outbound
|
||||
network requests are required during startup.
|
||||
"""
|
||||
# Use a minimal config that doesn't require external services
|
||||
minimal_config = """
|
||||
model_list:
|
||||
- model_name: fake-model
|
||||
litellm_params:
|
||||
model: fake/fake-model
|
||||
|
||||
general_settings:
|
||||
master_key: sk-test-1234
|
||||
database_url: null
|
||||
|
||||
environment_variables: {}
|
||||
"""
|
||||
|
||||
# Create a temporary config file
|
||||
config_path = "/tmp/litellm_test_config.yaml"
|
||||
with open(config_path, "w") as f:
|
||||
f.write(minimal_config)
|
||||
|
||||
image_to_use = "litellm/litellm"
|
||||
# Check if image exists
|
||||
result = subprocess.run(
|
||||
["docker", "images", "-q", image_to_use],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=10,
|
||||
)
|
||||
if not result.stdout.strip():
|
||||
pytest.skip(f"Docker image {image_to_use} not available")
|
||||
|
||||
# Run container with network disabled
|
||||
run_cmd = [
|
||||
"docker",
|
||||
"run",
|
||||
"--name",
|
||||
self.TEST_CONTAINER_NAME,
|
||||
"--network=none", # Disable all network access
|
||||
"-v",
|
||||
f"{config_path}:/app/config.yaml:ro",
|
||||
"-e",
|
||||
"LITELLM_MASTER_KEY=sk-test-1234",
|
||||
"-e",
|
||||
"DATABASE_URL=", # Empty to disable DB
|
||||
"-e",
|
||||
"STORE_MODEL_IN_DB=false",
|
||||
"-e",
|
||||
"LITELLM_LOG=DEBUG",
|
||||
"-d", # Detached mode
|
||||
image_to_use,
|
||||
"--config",
|
||||
"/app/config.yaml",
|
||||
]
|
||||
|
||||
result = subprocess.run(
|
||||
run_cmd,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=30,
|
||||
)
|
||||
|
||||
if result.returncode != 0:
|
||||
pytest.fail(
|
||||
f"Failed to start container: {result.stderr}"
|
||||
)
|
||||
|
||||
# Wait for container to start up
|
||||
time.sleep(5)
|
||||
|
||||
# Check container logs for any network-related errors
|
||||
logs_result = subprocess.run(
|
||||
["docker", "logs", self.TEST_CONTAINER_NAME],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=30,
|
||||
)
|
||||
|
||||
logs = logs_result.stdout + logs_result.stderr
|
||||
|
||||
# Check for network error patterns (case-insensitive)
|
||||
network_errors = []
|
||||
for pattern in self.NETWORK_ERROR_PATTERNS:
|
||||
matches = re.findall(pattern, logs, re.IGNORECASE)
|
||||
if matches:
|
||||
network_errors.extend(matches)
|
||||
|
||||
# Check if container is still running
|
||||
inspect_result = subprocess.run(
|
||||
[
|
||||
"docker",
|
||||
"inspect",
|
||||
"-f",
|
||||
"{{.State.Running}}",
|
||||
self.TEST_CONTAINER_NAME,
|
||||
],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=10,
|
||||
)
|
||||
|
||||
is_running = inspect_result.stdout.strip() == "true"
|
||||
|
||||
# Print diagnostic info
|
||||
print("\n=== Container Startup Logs (first 100 lines) ===")
|
||||
log_lines = logs.split("\n")[:100]
|
||||
for line in log_lines:
|
||||
print(line)
|
||||
print("=" * 50)
|
||||
|
||||
# If container crashed, get exit code and reason
|
||||
if not is_running:
|
||||
exit_result = subprocess.run(
|
||||
[
|
||||
"docker",
|
||||
"inspect",
|
||||
"-f",
|
||||
"{{.State.ExitCode}}",
|
||||
self.TEST_CONTAINER_NAME,
|
||||
],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=10,
|
||||
)
|
||||
exit_code = exit_result.stdout.strip()
|
||||
|
||||
# Container not running is OK if it didn't crash due to network issues
|
||||
# Check if exit was due to network errors
|
||||
if network_errors:
|
||||
pytest.fail(
|
||||
f"Container failed with network errors (exit code {exit_code}): "
|
||||
f"{network_errors}\n\nFull logs:\n{logs}"
|
||||
)
|
||||
|
||||
# Assert no network errors were found
|
||||
assert len(network_errors) == 0, (
|
||||
f"Container made network requests during startup that failed: "
|
||||
f"{network_errors}"
|
||||
)
|
||||
|
||||
def test_no_external_urls_in_startup_code(self):
|
||||
"""
|
||||
Static analysis test: check that startup code doesn't contain
|
||||
hardcoded external URLs that would be called during import/startup.
|
||||
|
||||
This is a complementary test to catch issues without needing Docker.
|
||||
"""
|
||||
# Directories to check for startup code
|
||||
startup_dirs = [
|
||||
"litellm/proxy",
|
||||
"litellm/__init__.py",
|
||||
"litellm/main.py",
|
||||
]
|
||||
|
||||
# Patterns that indicate external URL calls during startup (not in functions)
|
||||
problematic_patterns = [
|
||||
# Immediate HTTP calls (not inside functions)
|
||||
r'^requests\.get\(["\'](https?://)',
|
||||
r'^httpx\.get\(["\'](https?://)',
|
||||
r'^urllib\.request\.urlopen\(["\'](https?://)',
|
||||
]
|
||||
|
||||
# Files that are OK to have URLs (they're called on-demand, not startup)
|
||||
allowed_files = [
|
||||
"model_prices_and_context_window.json", # Static data file
|
||||
"test_", # Test files
|
||||
"_test.py",
|
||||
"conftest.py",
|
||||
]
|
||||
|
||||
workspace_root = os.path.dirname(
|
||||
os.path.dirname(os.path.dirname(os.path.dirname(__file__)))
|
||||
)
|
||||
|
||||
issues_found = []
|
||||
|
||||
for startup_dir in startup_dirs:
|
||||
full_path = os.path.join(workspace_root, startup_dir)
|
||||
if not os.path.exists(full_path):
|
||||
continue
|
||||
|
||||
if os.path.isfile(full_path):
|
||||
files_to_check = [full_path]
|
||||
else:
|
||||
files_to_check = []
|
||||
for root, _dirs, files in os.walk(full_path):
|
||||
for f in files:
|
||||
if f.endswith(".py"):
|
||||
files_to_check.append(os.path.join(root, f))
|
||||
|
||||
for filepath in files_to_check:
|
||||
# Skip allowed files
|
||||
if any(allowed in filepath for allowed in allowed_files):
|
||||
continue
|
||||
|
||||
try:
|
||||
with open(filepath, "r") as f:
|
||||
content = f.read()
|
||||
|
||||
for pattern in problematic_patterns:
|
||||
matches = re.findall(
|
||||
pattern, content, re.MULTILINE
|
||||
)
|
||||
if matches:
|
||||
issues_found.append(
|
||||
f"{filepath}: {pattern} matched"
|
||||
)
|
||||
except Exception:
|
||||
pass # Skip unreadable files
|
||||
|
||||
# This test is informational - we document but don't fail
|
||||
if issues_found:
|
||||
print(
|
||||
"\n=== Potential startup network calls found ===\n"
|
||||
+ "\n".join(issues_found)
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not is_docker_available(),
|
||||
reason="Docker not available",
|
||||
)
|
||||
def test_container_build_no_network_fetch():
|
||||
"""
|
||||
Test that the Docker build process doesn't require network for runtime.
|
||||
|
||||
This verifies that all dependencies are properly bundled and no
|
||||
runtime network calls are made during container initialization.
|
||||
|
||||
Note: Build itself may need network for pip install, but runtime should not.
|
||||
"""
|
||||
# This is a simplified version - full test would need to:
|
||||
# 1. Build image with --network=none (requires pre-cached deps)
|
||||
# 2. Or run built image in isolated network
|
||||
|
||||
# For now, just verify the Dockerfile doesn't have wget/curl in CMD/ENTRYPOINT
|
||||
workspace_root = os.path.dirname(
|
||||
os.path.dirname(os.path.dirname(os.path.dirname(__file__)))
|
||||
)
|
||||
dockerfile_path = os.path.join(workspace_root, "Dockerfile")
|
||||
|
||||
if not os.path.exists(dockerfile_path):
|
||||
pytest.skip("Dockerfile not found")
|
||||
|
||||
with open(dockerfile_path, "r") as f:
|
||||
content = f.read()
|
||||
|
||||
# Check for network calls in CMD/ENTRYPOINT
|
||||
problematic = []
|
||||
lines = content.split("\n")
|
||||
for i, line in enumerate(lines, 1):
|
||||
line_upper = line.strip().upper()
|
||||
if line_upper.startswith(("CMD", "ENTRYPOINT")):
|
||||
if any(
|
||||
cmd in line.lower()
|
||||
for cmd in ["curl", "wget", "fetch", "http://", "https://"]
|
||||
):
|
||||
problematic.append(f"Line {i}: {line.strip()}")
|
||||
|
||||
assert len(problematic) == 0, (
|
||||
f"Dockerfile CMD/ENTRYPOINT contains network calls: {problematic}"
|
||||
)
|
||||
|
||||
print("\nDockerfile CMD/ENTRYPOINT does not contain network calls.")
|
||||
Loading…
Add table
Reference in a new issue