mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
Enforce realtime resolved model scopes
This commit is contained in:
parent
b94a3061dd
commit
2961910849
2 changed files with 269 additions and 10 deletions
|
|
@ -3126,19 +3126,90 @@ async def can_key_call_resolved_model(
|
|||
valid_token: UserAPIKeyAuth,
|
||||
llm_router: Optional[litellm.Router],
|
||||
) -> None:
|
||||
if valid_token.config:
|
||||
return
|
||||
if (
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
skip_key_model_check = valid_token.config or (
|
||||
isinstance(valid_token.models, list)
|
||||
and SpecialModelNames.all_team_models.value in valid_token.models
|
||||
):
|
||||
return
|
||||
await can_key_call_model(
|
||||
model=model,
|
||||
llm_model_list=llm_model_list,
|
||||
valid_token=valid_token,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
if not skip_key_model_check:
|
||||
await can_key_call_model(
|
||||
model=model,
|
||||
llm_model_list=llm_model_list,
|
||||
valid_token=valid_token,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
|
||||
team_object: Optional[LiteLLM_TeamTableCachedObj] = None
|
||||
team_object_from_lookup = False
|
||||
if valid_token.team_id is not None:
|
||||
try:
|
||||
team_object = await get_team_object(
|
||||
team_id=valid_token.team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=valid_token.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
team_object_from_lookup = True
|
||||
except Exception:
|
||||
team_object = LiteLLM_TeamTableCachedObj(
|
||||
team_id=valid_token.team_id,
|
||||
models=valid_token.team_models,
|
||||
blocked=valid_token.team_blocked,
|
||||
team_alias=valid_token.team_alias,
|
||||
metadata=valid_token.team_metadata,
|
||||
object_permission_id=valid_token.team_object_permission_id,
|
||||
object_permission=valid_token.team_object_permission,
|
||||
)
|
||||
|
||||
if team_object is not None:
|
||||
try:
|
||||
await can_team_access_model(
|
||||
model=model,
|
||||
team_object=team_object,
|
||||
llm_router=llm_router,
|
||||
team_model_aliases=valid_token.team_model_aliases,
|
||||
)
|
||||
except ProxyException as team_denial:
|
||||
if team_denial.type != ProxyErrorTypes.team_model_access_denied:
|
||||
raise
|
||||
if not await _key_access_group_grants_model(
|
||||
model=model,
|
||||
valid_token=valid_token,
|
||||
team_object=team_object,
|
||||
llm_router=llm_router,
|
||||
):
|
||||
raise
|
||||
|
||||
if valid_token.user_id is not None and team_object_from_lookup:
|
||||
await _check_team_member_model_access(
|
||||
model=model,
|
||||
team_object=team_object,
|
||||
valid_token=valid_token,
|
||||
llm_router=llm_router,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
if valid_token.project_id is not None:
|
||||
project_object = await get_project_object(
|
||||
project_id=valid_token.project_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
if project_object is not None and len(project_object.models) > 0:
|
||||
can_project_access_model(
|
||||
model=model,
|
||||
project_object=project_object,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
|
||||
|
||||
def can_org_access_model(
|
||||
|
|
|
|||
|
|
@ -467,6 +467,152 @@ async def test_transcription_sessions_rejects_disallowed_resolved_model(
|
|||
proxy_app.dependency_overrides.pop(user_api_key_auth, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_transcription_sessions_rejects_disallowed_team_model_scope(
|
||||
proxy_app,
|
||||
):
|
||||
from litellm.proxy._types import LiteLLM_TeamTableCachedObj
|
||||
|
||||
team = LiteLLM_TeamTableCachedObj(
|
||||
team_id="team-a",
|
||||
models=["gpt-4o-realtime-preview"],
|
||||
)
|
||||
proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||||
user_id="test-user",
|
||||
team_id="team-a",
|
||||
models=["*"],
|
||||
)
|
||||
try:
|
||||
client = TestClient(proxy_app, raise_server_exceptions=False)
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.route_request") as mock_route_request,
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging,
|
||||
patch(
|
||||
"litellm.proxy.auth.auth_checks.get_team_object",
|
||||
new=AsyncMock(return_value=team),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.auth.auth_checks.get_team_membership",
|
||||
new=AsyncMock(return_value=None),
|
||||
),
|
||||
):
|
||||
mock_logging.post_call_failure_hook = AsyncMock()
|
||||
|
||||
response = client.post(
|
||||
"/v1/realtime/transcription_sessions",
|
||||
headers={"Authorization": "Bearer sk-test-master-key"},
|
||||
json={
|
||||
"input_audio_transcription": {"model": "gpt-realtime-whisper"}
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 403
|
||||
assert "team" in response.text.lower()
|
||||
assert "Tried to access gpt-realtime-whisper" in response.text
|
||||
mock_route_request.assert_not_called()
|
||||
finally:
|
||||
proxy_app.dependency_overrides.pop(user_api_key_auth, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_transcription_sessions_rejects_disallowed_project_model_scope(
|
||||
proxy_app,
|
||||
):
|
||||
from litellm.proxy._types import LiteLLM_ProjectTableCachedObj
|
||||
|
||||
project = LiteLLM_ProjectTableCachedObj(
|
||||
project_id="project-a",
|
||||
models=["gpt-4o-realtime-preview"],
|
||||
created_by="test-user",
|
||||
updated_by="test-user",
|
||||
)
|
||||
proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||||
user_id="test-user",
|
||||
project_id="project-a",
|
||||
models=["*"],
|
||||
)
|
||||
try:
|
||||
client = TestClient(proxy_app, raise_server_exceptions=False)
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.route_request") as mock_route_request,
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging,
|
||||
patch(
|
||||
"litellm.proxy.auth.auth_checks.get_project_object",
|
||||
new=AsyncMock(return_value=project),
|
||||
),
|
||||
):
|
||||
mock_logging.post_call_failure_hook = AsyncMock()
|
||||
|
||||
response = client.post(
|
||||
"/v1/realtime/transcription_sessions",
|
||||
headers={"Authorization": "Bearer sk-test-master-key"},
|
||||
json={
|
||||
"input_audio_transcription": {"model": "gpt-realtime-whisper"}
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 403
|
||||
assert "project" in response.text.lower()
|
||||
assert "Tried to access gpt-realtime-whisper" in response.text
|
||||
mock_route_request.assert_not_called()
|
||||
finally:
|
||||
proxy_app.dependency_overrides.pop(user_api_key_auth, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_transcription_sessions_rejects_disallowed_team_member_model_scope(
|
||||
proxy_app,
|
||||
):
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_BudgetTable,
|
||||
LiteLLM_TeamMembership,
|
||||
LiteLLM_TeamTableCachedObj,
|
||||
)
|
||||
|
||||
team = LiteLLM_TeamTableCachedObj(team_id="team-a", models=["*"])
|
||||
membership = LiteLLM_TeamMembership(
|
||||
user_id="test-user",
|
||||
team_id="team-a",
|
||||
litellm_budget_table=LiteLLM_BudgetTable(
|
||||
allowed_models=["gpt-4o-realtime-preview"],
|
||||
),
|
||||
)
|
||||
proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||||
user_id="test-user",
|
||||
team_id="team-a",
|
||||
models=["*"],
|
||||
)
|
||||
try:
|
||||
client = TestClient(proxy_app, raise_server_exceptions=False)
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.route_request") as mock_route_request,
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging,
|
||||
patch(
|
||||
"litellm.proxy.auth.auth_checks.get_team_object",
|
||||
new=AsyncMock(return_value=team),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.auth.auth_checks.get_team_membership",
|
||||
new=AsyncMock(return_value=membership),
|
||||
),
|
||||
):
|
||||
mock_logging.post_call_failure_hook = AsyncMock()
|
||||
|
||||
response = client.post(
|
||||
"/v1/realtime/transcription_sessions",
|
||||
headers={"Authorization": "Bearer sk-test-master-key"},
|
||||
json={
|
||||
"input_audio_transcription": {"model": "gpt-realtime-whisper"}
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 403
|
||||
assert "Team member not allowed to access model" in response.text
|
||||
mock_route_request.assert_not_called()
|
||||
finally:
|
||||
proxy_app.dependency_overrides.pop(user_api_key_auth, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_realtime_transcription_websocket_default_model_checks_key_scope():
|
||||
from litellm.proxy import proxy_server
|
||||
|
|
@ -490,6 +636,48 @@ async def test_realtime_transcription_websocket_default_model_checks_key_scope()
|
|||
assert "not allowed to access model" in close_kwargs["reason"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_realtime_transcription_websocket_default_model_checks_team_scope():
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._types import LiteLLM_TeamTableCachedObj
|
||||
|
||||
team = LiteLLM_TeamTableCachedObj(
|
||||
team_id="team-a",
|
||||
models=["gpt-4o-realtime-preview"],
|
||||
)
|
||||
websocket = MagicMock()
|
||||
websocket.headers = {}
|
||||
websocket.close = AsyncMock()
|
||||
websocket.accept = AsyncMock()
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.auth.auth_checks.get_team_object",
|
||||
new=AsyncMock(return_value=team),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.auth.auth_checks.get_team_membership",
|
||||
new=AsyncMock(return_value=None),
|
||||
),
|
||||
):
|
||||
await proxy_server.realtime_websocket_endpoint(
|
||||
websocket=websocket,
|
||||
model=None,
|
||||
intent="transcription",
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_id="test-user",
|
||||
team_id="team-a",
|
||||
models=["*"],
|
||||
),
|
||||
)
|
||||
|
||||
websocket.accept.assert_not_awaited()
|
||||
websocket.close.assert_awaited_once()
|
||||
_, close_kwargs = websocket.close.call_args
|
||||
assert close_kwargs["code"] == 1008
|
||||
assert "not allowed to access model" in close_kwargs["reason"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_transcription_sessions_encrypts_client_secret(
|
||||
proxy_app,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue