fix(test): make test_end_user_jwt_auth resilient to module reloads (#24480)

test_cors_config does sys.modules.pop("litellm.proxy.proxy_server") +
re-import, creating a new module object.  The JWT test's setattr calls
were writing to the old module while runtime code read from the new one,
causing ProxyModelNotFoundError in CI.

Fix by patching both the original module's __globals__ dict (used by
chat_completion) and the live sys.modules entry (used by lazy imports in
user_api_key_auth).

Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
Aarish Alam 2026-03-24 12:00:55 +05:30 committed by GitHub
parent 975dfe3fa8
commit f414b72ad4
No known key found for this signature in database
GPG key ID: B5690EEEBB952194

View file

@ -1170,7 +1170,7 @@ async def test_end_user_jwt_auth(monkeypatch):
# use generated key to auth in
from litellm import Router
from litellm.types.router import RouterGeneralSettings
# Create a router with pass_through_all_models enabled
router = Router(
model_list=[],
@ -1178,75 +1178,91 @@ async def test_end_user_jwt_auth(monkeypatch):
pass_through_all_models=True
),
)
setattr(litellm.proxy.proxy_server, "premium_user", True)
setattr(
litellm.proxy.proxy_server,
"general_settings",
{"enable_jwt_auth": True},
)
setattr(
litellm.proxy.proxy_server,
"llm_router",
router,
)
setattr(litellm.proxy.proxy_server, "prisma_client", {})
setattr(litellm.proxy.proxy_server, "jwt_handler", jwt_handler)
from litellm.proxy.proxy_server import cost_tracking
cost_tracking()
result = await user_api_key_auth(request=request, api_key=bearer_token)
# Assert that end_user_id is correctly extracted from JWT token's 'sub' field
assert result.end_user_id == "81b3e52a-67a6-4efb-9645-70527e101479"
# Other tests (e.g. test_cors_config) do sys.modules.pop + re-import of
# litellm.proxy.proxy_server, creating a new module object. Functions
# captured before the reload (like chat_completion at the top of this file)
# read globals from the *original* module dict via __globals__, while lazy
# imports inside user_api_key_auth read from whatever is currently in
# sys.modules. Patch both dicts so all code paths see our values.
_attrs = {
"premium_user": True,
"general_settings": {"enable_jwt_auth": True},
"llm_router": router,
"prisma_client": {},
"jwt_handler": jwt_handler,
}
temp_response = Response()
from litellm.proxy.hooks.proxy_track_cost_callback import (
_should_track_cost_callback,
)
_g = chat_completion.__globals__
_live_mod = sys.modules.get("litellm.proxy.proxy_server")
_dicts_to_patch = [_g]
if _live_mod is not None and _live_mod.__dict__ is not _g:
_dicts_to_patch.append(_live_mod.__dict__)
# Mock the actual LLM completion call
mock_response = litellm.ModelResponse(
id="chatcmpl-mock",
choices=[
litellm.Choices(
finish_reason="stop",
index=0,
message=litellm.Message(
content="Hello! I'm doing well, thank you for asking.",
role="assistant",
),
)
],
created=1234567890,
model="gpt-4o",
object="chat.completion",
usage=litellm.Usage(
prompt_tokens=10,
completion_tokens=15,
total_tokens=25,
),
)
_saved = []
for _d in _dicts_to_patch:
_saved.append({k: _d.get(k) for k in _attrs})
_d.update(_attrs)
with patch("litellm.acompletion", new=AsyncMock(return_value=mock_response)) as mock_completion:
resp = await chat_completion(
request=request,
fastapi_response=temp_response,
model="gpt-4o",
user_api_key_dict=result,
try:
from litellm.proxy.proxy_server import cost_tracking
cost_tracking()
result = await user_api_key_auth(request=request, api_key=bearer_token)
# Assert that end_user_id is correctly extracted from JWT token's 'sub' field
assert result.end_user_id == "81b3e52a-67a6-4efb-9645-70527e101479"
temp_response = Response()
from litellm.proxy.hooks.proxy_track_cost_callback import (
_should_track_cost_callback,
)
assert resp is not None
# Mock the actual LLM completion call
mock_response = litellm.ModelResponse(
id="chatcmpl-mock",
choices=[
litellm.Choices(
finish_reason="stop",
index=0,
message=litellm.Message(
content="Hello! I'm doing well, thank you for asking.",
role="assistant",
),
)
],
created=1234567890,
model="gpt-4o",
object="chat.completion",
usage=litellm.Usage(
prompt_tokens=10,
completion_tokens=15,
total_tokens=25,
),
)
await asyncio.sleep(1)
with patch("litellm.acompletion", new=AsyncMock(return_value=mock_response)) as mock_completion:
resp = await chat_completion(
request=request,
fastapi_response=temp_response,
model="gpt-4o",
user_api_key_dict=result,
)
# Verify the completion was called with correct end_user_id
mock_completion.assert_called_once()
call_kwargs = mock_completion.call_args.kwargs
# end_user_id is passed in metadata as 'user_api_key_end_user_id'
metadata = call_kwargs.get("metadata", {})
assert metadata.get("user_api_key_end_user_id") == "81b3e52a-67a6-4efb-9645-70527e101479"
assert resp is not None
await asyncio.sleep(1)
# Verify the completion was called with correct end_user_id
mock_completion.assert_called_once()
call_kwargs = mock_completion.call_args.kwargs
# end_user_id is passed in metadata as 'user_api_key_end_user_id'
metadata = call_kwargs.get("metadata", {})
assert metadata.get("user_api_key_end_user_id") == "81b3e52a-67a6-4efb-9645-70527e101479"
finally:
for _d, _s in zip(_dicts_to_patch, _saved):
_d.update(_s)
def test_can_rbac_role_call_route():