mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Merge branch 'litellm_internal_staging' into litellm_standardize_rate_limit_errors-5fb4
This commit is contained in:
commit
ff2c03b3d3
8 changed files with 272 additions and 21 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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';
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
}
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
);
|
||||
|
||||
|
|
|
|||
|
|
@ -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?'}]}]};
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
);
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
# ============================================
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue