Merge branch 'litellm_internal_staging' into litellm_standardize_rate_limit_errors-5fb4

This commit is contained in:
mateo-berri 2026-06-03 17:33:16 +00:00
commit ff2c03b3d3
No known key found for this signature in database
8 changed files with 272 additions and 21 deletions

View file

@ -894,7 +894,12 @@ async def _common_key_generation_helper( # noqa: PLR0915
user_api_key_dict.user_role is not None
and user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value
)
if not _is_proxy_admin:
_org_inherited_from_team = (
team_table is not None
and team_table.organization_id is not None
and data.organization_id == team_table.organization_id
)
if not _is_proxy_admin and not _org_inherited_from_team:
await _validate_caller_can_assign_key_org(
user_api_key_dict=user_api_key_dict,
organization_id=data.organization_id,

View file

@ -32,7 +32,7 @@ describe('Gemini AI Tests', () => {
};
const model = genAI.getGenerativeModel({
model: 'gemini-2.5-flash-lite'
model: 'gemini-3.1-flash-lite'
}, requestOptions);
const prompt = 'Say "hello test" and nothing else';
@ -83,7 +83,7 @@ describe('Gemini AI Tests', () => {
};
const model = genAI.getGenerativeModel({
model: 'gemini-2.5-flash-lite'
model: 'gemini-3.1-flash-lite'
}, requestOptions);
const prompt = 'Say "hello test" and nothing else';

View file

@ -1,13 +1,13 @@
const { GoogleGenerativeAI, ModelParams, RequestOptions } = require("@google/generative-ai");
const modelParams = {
model: 'gemini-2.5-flash-lite',
model: 'gemini-3.1-flash-lite',
};
const requestOptions = {
baseUrl: 'http://127.0.0.1:4000/gemini',
customHeaders: {
"tags": "gemini-js-sdk,gemini-2.5-flash-lite"
"tags": "gemini-js-sdk,gemini-3.1-flash-lite"
}
};

View file

@ -4,7 +4,7 @@ const { VertexAI, RequestOptions } = require('@google-cloud/vertexai');
const vertexAI = new VertexAI({
project: 'litellm-ci-cd',
location: 'us-central1',
location: 'global',
apiEndpoint: "127.0.0.1:4000/vertex-ai"
});
@ -20,7 +20,7 @@ const requestOptions = {
};
const generativeModel = vertexAI.getGenerativeModel(
{ model: 'gemini-2.5-flash-lite' },
{ model: 'gemini-3.1-flash-lite' },
requestOptions
);

View file

@ -56,6 +56,9 @@ beforeAll(() => {
loadVertexAiCredentials();
});
// Configure Jest to retry flaky tests up to 3 times (useful for 429 rate limiting)
jest.retryTimes(3);
// Non-streaming Vertex generateContent can exceed 5s in CI / under load
const VERTEX_TEST_TIMEOUT_MS = 30000;
@ -65,7 +68,7 @@ describe('Vertex AI Tests', () => {
async () => {
const vertexAI = new VertexAI({
project: 'litellm-ci-cd',
location: 'us-central1',
location: 'global',
apiEndpoint: "localhost:4000/vertex-ai"
});
@ -78,7 +81,7 @@ describe('Vertex AI Tests', () => {
};
const generativeModel = vertexAI.getGenerativeModel(
{ model: 'gemini-2.5-flash-lite' },
{ model: 'gemini-3.1-flash-lite' },
requestOptions
);
@ -108,13 +111,13 @@ describe('Vertex AI Tests', () => {
async () => {
const vertexAI = new VertexAI({
project: 'litellm-ci-cd',
location: 'us-central1',
location: 'global',
apiEndpoint: "localhost:4000/vertex-ai"
});
const customHeaders = new Headers({"x-litellm-api-key": "sk-1234"});
const requestOptions = {customHeaders: customHeaders};
const generativeModel = vertexAI.getGenerativeModel(
{model: 'gemini-2.5-flash-lite'},
{model: 'gemini-3.1-flash-lite'},
requestOptions
);
const request = {contents: [{role: 'user', parts: [{text: 'What is 2+2?'}]}]};

View file

@ -103,12 +103,12 @@ async def test_basic_vertex_ai_pass_through_with_spendlog():
vertexai.init(
project="litellm-ci-cd",
location="us-central1",
location="global",
api_endpoint=f"{LITE_LLM_ENDPOINT}/vertex_ai",
api_transport="rest",
)
model = GenerativeModel(model_name="gemini-2.5-flash-lite")
model = GenerativeModel(model_name="gemini-3.1-flash-lite")
response = model.generate_content("hi")
print("response", response)
@ -143,12 +143,12 @@ async def test_basic_vertex_ai_pass_through_streaming_with_spendlog():
vertexai.init(
project="litellm-ci-cd",
location="us-central1",
location="global",
api_endpoint=f"{LITE_LLM_ENDPOINT}/vertex_ai",
api_transport="rest",
)
model = GenerativeModel(model_name="gemini-2.5-flash-lite")
model = GenerativeModel(model_name="gemini-3.1-flash-lite")
response = model.generate_content("hi", stream=True)
for chunk in response:
@ -182,7 +182,7 @@ async def test_vertex_ai_pass_through_endpoint_context_caching():
vertexai.init(
project="litellm-ci-cd",
location="us-central1",
location="global",
api_endpoint=f"{LITE_LLM_ENDPOINT}/vertex_ai",
api_transport="rest",
)
@ -204,7 +204,7 @@ async def test_vertex_ai_pass_through_endpoint_context_caching():
]
cached_content = caching.CachedContent.create(
model_name="gemini-2.5-flash-lite-001",
model_name="gemini-3.1-flash-lite",
system_instruction=system_instruction,
contents=contents,
ttl=datetime.timedelta(minutes=60),

View file

@ -71,7 +71,7 @@ describe('Vertex AI Tests', () => {
test('should successfully generate non-streaming content with tags', async () => {
const vertexAI = new VertexAI({
project: 'litellm-ci-cd',
location: 'us-central1',
location: 'global',
apiEndpoint: "127.0.0.1:4000/vertex_ai"
});
@ -85,7 +85,7 @@ describe('Vertex AI Tests', () => {
};
const generativeModel = vertexAI.getGenerativeModel(
{ model: 'gemini-2.5-flash-lite' },
{ model: 'gemini-3.1-flash-lite' },
requestOptions
);
@ -130,7 +130,7 @@ describe('Vertex AI Tests', () => {
test('should successfully generate streaming content with tags', async () => {
const vertexAI = new VertexAI({
project: 'litellm-ci-cd',
location: 'us-central1',
location: 'global',
apiEndpoint: "127.0.0.1:4000/vertex_ai"
});
@ -144,7 +144,7 @@ describe('Vertex AI Tests', () => {
};
const generativeModel = vertexAI.getGenerativeModel(
{ model: 'gemini-2.5-flash-lite' },
{ model: 'gemini-3.1-flash-lite' },
requestOptions
);

View file

@ -3094,6 +3094,249 @@ async def test_generate_key_with_object_permission():
assert "object_permission" not in key_data
@pytest.mark.asyncio
async def test_generate_key_team_member_inherits_org_skips_membership_check():
"""Regression: a team member creating a key for an org-scoped team must not
be blocked by the org-membership check.
When ``organization_id`` is inherited from the key's team (via
``apply_enterprise_key_management_params`` -> ``add_team_organization_id``),
the caller already passed team-level authorization. Requiring an explicit
``LiteLLM_OrganizationMembership`` row on top of that broke the normal admin
workflow (admins only add users to teams). This asserts the org-membership
check is skipped when the org id came from the caller's team.
"""
from unittest.mock import AsyncMock, MagicMock, patch
from litellm.proxy._types import GenerateKeyRequest, LitellmUserRoles
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth
from litellm.proxy.management_endpoints.key_management_endpoints import (
_common_key_generation_helper,
)
org_id = "org-from-team"
# Team belongs to an org; caller is a team member but NOT an explicit member
# of that organization (the regression scenario).
mock_team_table = MagicMock()
mock_team_table.organization_id = org_id
mock_team_table.metadata = None
mock_validate_org = AsyncMock()
mock_generate_key = AsyncMock(
return_value={
"key": "sk-test-key",
"expires": None,
"user_id": "alice",
"team_id": "team-1",
}
)
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.llm_router", None),
patch("litellm.proxy.proxy_server.premium_user", True),
patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.validate_key_mcp_servers_against_team",
new_callable=AsyncMock,
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.validate_key_search_tools_against_team",
new_callable=AsyncMock,
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._validate_caller_can_assign_key_org",
mock_validate_org,
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.get_org_object",
new_callable=AsyncMock,
return_value=MagicMock(litellm_budget_table=None),
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._check_org_key_limits",
new_callable=AsyncMock,
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn",
mock_generate_key,
),
):
result = await _common_key_generation_helper(
data=GenerateKeyRequest(
user_id="alice",
team_id="team-1",
organization_id=org_id,
),
user_api_key_dict=UserAPIKeyAuth(
user_id="alice",
user_role=LitellmUserRoles.INTERNAL_USER.value,
),
litellm_changed_by=None,
team_table=mock_team_table,
)
# Key creation proceeded for the team member ...
mock_generate_key.assert_awaited_once()
assert result is not None
# ... and the org-membership check was bypassed because organization_id was
# inherited from the caller's team.
mock_validate_org.assert_not_awaited()
@pytest.mark.asyncio
async def test_generate_key_foreign_org_without_team_still_enforces_membership():
"""VERIA-55: a caller assigning a key to an organization that was NOT
inherited from a team must still pass the org-membership check.
This guards the IDOR fix: ``team_table is None`` (or an org id that does not
match the team) means the org id did not come from team context, so the
explicit membership validation must run.
"""
from unittest.mock import AsyncMock, MagicMock, patch
from litellm.proxy._types import GenerateKeyRequest, LitellmUserRoles
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth
from litellm.proxy.management_endpoints.key_management_endpoints import (
_common_key_generation_helper,
)
foreign_org_id = "someone-elses-org"
mock_validate_org = AsyncMock()
mock_generate_key = AsyncMock(
return_value={
"key": "sk-test-key",
"expires": None,
"user_id": "alice",
"team_id": None,
}
)
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.llm_router", None),
patch("litellm.proxy.proxy_server.premium_user", True),
patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._validate_caller_can_assign_key_org",
mock_validate_org,
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.get_org_object",
new_callable=AsyncMock,
return_value=MagicMock(litellm_budget_table=None),
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._check_org_key_limits",
new_callable=AsyncMock,
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn",
mock_generate_key,
),
):
await _common_key_generation_helper(
data=GenerateKeyRequest(
user_id="alice",
organization_id=foreign_org_id,
),
user_api_key_dict=UserAPIKeyAuth(
user_id="alice",
user_role=LitellmUserRoles.INTERNAL_USER.value,
),
litellm_changed_by=None,
team_table=None,
)
# No team context -> the org-membership check must still run.
mock_validate_org.assert_awaited_once()
assert mock_validate_org.call_args.kwargs["organization_id"] == foreign_org_id
@pytest.mark.asyncio
async def test_generate_key_foreign_org_with_mismatched_team_still_enforces_membership():
"""VERIA-55: when a team is present but its organization_id differs from the
organization_id on the key request, the org-membership check must still run."""
from unittest.mock import AsyncMock, MagicMock, patch
from litellm.proxy._types import GenerateKeyRequest, LitellmUserRoles
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth
from litellm.proxy.management_endpoints.key_management_endpoints import (
_common_key_generation_helper,
)
team_org_id = "other-org"
foreign_org_id = "someone-elses-org"
mock_team_table = MagicMock()
mock_team_table.organization_id = team_org_id
mock_team_table.metadata = None
mock_validate_org = AsyncMock()
mock_generate_key = AsyncMock(
return_value={
"key": "sk-test-key",
"expires": None,
"user_id": "alice",
"team_id": "team-1",
}
)
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.llm_router", None),
patch("litellm.proxy.proxy_server.premium_user", True),
patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.validate_key_mcp_servers_against_team",
new_callable=AsyncMock,
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.validate_key_search_tools_against_team",
new_callable=AsyncMock,
),
patch(
"litellm_enterprise.proxy.management_endpoints.key_management_endpoints.apply_enterprise_key_management_params",
side_effect=lambda data, team_table: data,
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._validate_caller_can_assign_key_org",
mock_validate_org,
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.get_org_object",
new_callable=AsyncMock,
return_value=MagicMock(litellm_budget_table=None),
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._check_org_key_limits",
new_callable=AsyncMock,
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn",
mock_generate_key,
),
):
await _common_key_generation_helper(
data=GenerateKeyRequest(
user_id="alice",
team_id="team-1",
organization_id=foreign_org_id,
),
user_api_key_dict=UserAPIKeyAuth(
user_id="alice",
user_role=LitellmUserRoles.INTERNAL_USER.value,
),
litellm_changed_by=None,
team_table=mock_team_table,
)
mock_validate_org.assert_awaited_once()
assert mock_validate_org.call_args.kwargs["organization_id"] == foreign_org_id
# ============================================
# Organization Key Limit Tests
# ============================================