mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge pull request #18675 from BerriAI/litellm_key_team_routing_config
[Feature] Allow Add and Update for router_settings on Keys and Teams
This commit is contained in:
commit
e285e2b91d
14 changed files with 371 additions and 11 deletions
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.4.20-py3-none-any.whl
vendored
Normal file
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.4.20-py3-none-any.whl
vendored
Normal file
Binary file not shown.
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.4.20.tar.gz
vendored
Normal file
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.4.20.tar.gz
vendored
Normal file
Binary file not shown.
|
|
@ -0,0 +1,6 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_TeamTable" ADD COLUMN "router_settings" JSONB DEFAULT '{}';
|
||||
|
||||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_VerificationToken" ADD COLUMN "router_settings" JSONB DEFAULT '{}';
|
||||
|
||||
|
|
@ -124,6 +124,7 @@ model LiteLLM_TeamTable {
|
|||
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
|
||||
model_spend Json @default("{}")
|
||||
model_max_budget Json @default("{}")
|
||||
router_settings Json? @default("{}")
|
||||
team_member_permissions String[] @default([])
|
||||
model_id Int? @unique // id for LiteLLM_ModelTable -> stores team-level model aliases
|
||||
litellm_organization_table LiteLLM_OrganizationTable? @relation(fields: [organization_id], references: [organization_id])
|
||||
|
|
@ -225,6 +226,7 @@ model LiteLLM_VerificationToken {
|
|||
models String[]
|
||||
aliases Json @default("{}")
|
||||
config Json @default("{}")
|
||||
router_settings Json? @default("{}")
|
||||
user_id String?
|
||||
team_id String?
|
||||
permissions Json @default("{}")
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[tool.poetry]
|
||||
name = "litellm-proxy-extras"
|
||||
version = "0.4.19"
|
||||
version = "0.4.20"
|
||||
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
|
||||
authors = ["BerriAI"]
|
||||
readme = "README.md"
|
||||
|
|
@ -22,7 +22,7 @@ requires = ["poetry-core"]
|
|||
build-backend = "poetry.core.masonry.api"
|
||||
|
||||
[tool.commitizen]
|
||||
version = "0.4.19"
|
||||
version = "0.4.20"
|
||||
version_files = [
|
||||
"pyproject.toml:version",
|
||||
"../requirements.txt:litellm-proxy-extras==",
|
||||
|
|
|
|||
|
|
@ -863,6 +863,7 @@ class KeyRequestBase(GenerateRequestBase):
|
|||
tpm_limit_type: Optional[
|
||||
Literal["guaranteed_throughput", "best_effort_throughput", "dynamic"]
|
||||
] = None # raise an error if 'guaranteed_throughput' is set and we're overallocating tpm
|
||||
router_settings: Optional[UpdateRouterConfig] = None
|
||||
|
||||
|
||||
class LiteLLMKeyType(str, enum.Enum):
|
||||
|
|
@ -918,6 +919,7 @@ class GenerateKeyResponse(KeyRequestBase):
|
|||
"config",
|
||||
"permissions",
|
||||
"model_max_budget",
|
||||
"router_settings",
|
||||
]
|
||||
for field in dict_fields:
|
||||
value = values.get(field)
|
||||
|
|
@ -1460,6 +1462,7 @@ class TeamBase(LiteLLMPydanticObjectBase):
|
|||
|
||||
models: list = []
|
||||
blocked: bool = False
|
||||
router_settings: Optional[dict] = None
|
||||
|
||||
|
||||
class NewTeamRequest(TeamBase):
|
||||
|
|
@ -1542,6 +1545,7 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase):
|
|||
model_rpm_limit: Optional[Dict[str, int]] = None
|
||||
model_tpm_limit: Optional[Dict[str, int]] = None
|
||||
allowed_vector_store_indexes: Optional[List[AllowedVectorStoreIndexItem]] = None
|
||||
router_settings: Optional[dict] = None
|
||||
|
||||
|
||||
class ResetTeamBudgetRequest(LiteLLMPydanticObjectBase):
|
||||
|
|
@ -1684,6 +1688,7 @@ class LiteLLM_TeamTable(TeamBase):
|
|||
"permissions",
|
||||
"model_max_budget",
|
||||
"model_aliases",
|
||||
"router_settings",
|
||||
]
|
||||
|
||||
if isinstance(values, BaseModel):
|
||||
|
|
|
|||
|
|
@ -14,9 +14,10 @@ import copy
|
|||
import json
|
||||
import secrets
|
||||
import traceback
|
||||
import yaml
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import List, Literal, Optional, Tuple, cast
|
||||
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
import fastapi
|
||||
from fastapi import APIRouter, Depends, Header, HTTPException, Query, Request, status
|
||||
|
||||
|
|
@ -1033,7 +1034,7 @@ async def generate_key_fn(
|
|||
- auto_rotate: Optional[bool] - Whether this key should be automatically rotated (regenerated)
|
||||
- rotation_interval: Optional[str] - How often to auto-rotate this key (e.g., '30s', '30m', '30h', '30d'). Required if auto_rotate=True.
|
||||
- allowed_vector_store_indexes: Optional[List[dict]] - List of allowed vector store indexes for the key. Example - [{"index_name": "my-index", "index_permissions": ["write", "read"]}]. If specified, the key will only be able to use these specific vector store indexes. Create index, using `/v1/indexes` endpoint.
|
||||
|
||||
- router_settings: Optional[UpdateRouterConfig] - key-specific router settings. Example - {"model_group_retry_policy": {"max_retries": 5}}. IF null or {} then no router settings.
|
||||
|
||||
Examples:
|
||||
|
||||
|
|
@ -1388,6 +1389,10 @@ async def prepare_key_update_data(
|
|||
if "model_max_budget" in non_default_values:
|
||||
validate_model_max_budget(non_default_values["model_max_budget"])
|
||||
|
||||
# Serialize router_settings to JSON if present
|
||||
if "router_settings" in non_default_values and non_default_values["router_settings"] is not None:
|
||||
non_default_values["router_settings"] = safe_dumps(non_default_values["router_settings"])
|
||||
|
||||
non_default_values = prepare_metadata_fields(
|
||||
data=data, non_default_values=non_default_values, existing_metadata=_metadata
|
||||
)
|
||||
|
|
@ -1489,7 +1494,8 @@ async def update_key_fn(
|
|||
- auto_rotate: Optional[bool] - Whether this key should be automatically rotated
|
||||
- rotation_interval: Optional[str] - How often to rotate this key (e.g., '30d', '90d'). Required if auto_rotate=True
|
||||
- allowed_vector_store_indexes: Optional[List[dict]] - List of allowed vector store indexes for the key. Example - [{"index_name": "my-index", "index_permissions": ["write", "read"]}]. If specified, the key will only be able to use these specific vector store indexes. Create index, using `/v1/indexes` endpoint.
|
||||
|
||||
- router_settings: Optional[UpdateRouterConfig] - key-specific router settings. Example - {"model_group_retry_policy": {"max_retries": 5}}. IF null or {} then no router settings.
|
||||
|
||||
Example:
|
||||
```bash
|
||||
curl --location 'http://0.0.0.0:4000/key/update' \
|
||||
|
|
@ -2080,6 +2086,7 @@ async def generate_key_helper_fn( # noqa: PLR0915
|
|||
object_permission: Optional[LiteLLM_ObjectPermissionBase] = None,
|
||||
auto_rotate: Optional[bool] = None,
|
||||
rotation_interval: Optional[str] = None,
|
||||
router_settings: Optional[dict] = None,
|
||||
):
|
||||
from litellm.proxy.proxy_server import premium_user, prisma_client
|
||||
|
||||
|
|
@ -2114,6 +2121,7 @@ async def generate_key_helper_fn( # noqa: PLR0915
|
|||
aliases_json = json.dumps(aliases)
|
||||
config_json = json.dumps(config)
|
||||
permissions_json = json.dumps(permissions)
|
||||
router_settings_json = safe_dumps(router_settings) if router_settings is not None else safe_dumps({})
|
||||
|
||||
# Add model_rpm_limit and model_tpm_limit to metadata
|
||||
if model_rpm_limit is not None:
|
||||
|
|
@ -2189,6 +2197,7 @@ async def generate_key_helper_fn( # noqa: PLR0915
|
|||
"updated_by": updated_by,
|
||||
"allowed_routes": allowed_routes or [],
|
||||
"object_permission_id": object_permission_id,
|
||||
"router_settings": router_settings_json,
|
||||
}
|
||||
|
||||
# Add rotation fields if auto_rotate is enabled
|
||||
|
|
@ -2225,6 +2234,13 @@ async def generate_key_helper_fn( # noqa: PLR0915
|
|||
saved_token["model_max_budget"] = json.loads(
|
||||
saved_token["model_max_budget"]
|
||||
)
|
||||
router_settings = cast(Optional[dict], saved_token.get("router_settings"))
|
||||
if router_settings is not None and isinstance(router_settings, str):
|
||||
try:
|
||||
saved_token["router_settings"] = yaml.safe_load(router_settings)
|
||||
except yaml.YAMLError:
|
||||
# If it's not valid JSON/YAML, keep as is or set to empty dict
|
||||
saved_token["router_settings"] = {}
|
||||
|
||||
if saved_token.get("expires", None) is not None and isinstance(
|
||||
saved_token["expires"], datetime
|
||||
|
|
@ -2269,6 +2285,15 @@ async def generate_key_helper_fn( # noqa: PLR0915
|
|||
)
|
||||
key_data["created_at"] = getattr(create_key_response, "created_at", None)
|
||||
key_data["updated_at"] = getattr(create_key_response, "updated_at", None)
|
||||
|
||||
# Deserialize router_settings from JSON string to dict for response
|
||||
router_settings_value = key_data.get("router_settings")
|
||||
if router_settings_value is not None and isinstance(router_settings_value, str):
|
||||
try:
|
||||
key_data["router_settings"] = yaml.safe_load(router_settings_value)
|
||||
except yaml.YAMLError:
|
||||
# If it's not valid JSON/YAML, keep as is or set to empty dict
|
||||
key_data["router_settings"] = {}
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(
|
||||
"litellm.proxy.proxy_server.generate_key_helper_fn(): Exception occured - {}".format(
|
||||
|
|
|
|||
|
|
@ -100,7 +100,7 @@ from litellm.types.proxy.management_endpoints.team_endpoints import (
|
|||
TeamMemberAddResult,
|
||||
UpdateTeamMemberPermissionsRequest,
|
||||
)
|
||||
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
|
|
@ -696,8 +696,7 @@ async def new_team( # noqa: PLR0915
|
|||
- allowed_passthrough_routes: Optional[List[str]] - List of allowed pass through routes for the team.
|
||||
- allowed_vector_store_indexes: Optional[List[dict]] - List of allowed vector store indexes for the key. Example - [{"index_name": "my-index", "index_permissions": ["write", "read"]}]. If specified, the key will only be able to use these specific vector store indexes. Create index, using `/v1/indexes` endpoint.
|
||||
- secret_manager_settings: Optional[dict] - Secret manager settings for the team. [Docs](https://docs.litellm.ai/docs/secret_managers/overview)
|
||||
|
||||
|
||||
- router_settings: Optional[UpdateRouterConfig] - team-specific router settings. Example - {"model_group_retry_policy": {"max_retries": 5}}. IF null or {} then no router settings.
|
||||
|
||||
Returns:
|
||||
- team_id: (str) Unique team id - used for tracking spend across multiple keys for same team id.
|
||||
|
|
@ -911,6 +910,12 @@ async def new_team( # noqa: PLR0915
|
|||
complete_team_data.members_with_roles = []
|
||||
|
||||
complete_team_data_dict = complete_team_data.model_dump(exclude_none=True)
|
||||
|
||||
# Serialize router_settings to JSON (matching key creation pattern)
|
||||
router_settings_value = getattr(data, "router_settings", None)
|
||||
router_settings_json = safe_dumps(router_settings_value) if router_settings_value is not None else safe_dumps({})
|
||||
complete_team_data_dict["router_settings"] = router_settings_json
|
||||
|
||||
complete_team_data_dict = prisma_client.jsonify_team_object(
|
||||
db_data=complete_team_data_dict
|
||||
)
|
||||
|
|
@ -1234,7 +1239,7 @@ async def update_team( # noqa: PLR0915
|
|||
Example - update team TPM Limit
|
||||
- allowed_vector_store_indexes: Optional[List[dict]] - List of allowed vector store indexes for the key. Example - [{"index_name": "my-index", "index_permissions": ["write", "read"]}]. If specified, the key will only be able to use these specific vector store indexes. Create index, using `/v1/indexes` endpoint.
|
||||
- secret_manager_settings: Optional[dict] - Secret manager settings for the team. [Docs](https://docs.litellm.ai/docs/secret_managers/overview)
|
||||
|
||||
- router_settings: Optional[UpdateRouterConfig] - team-specific router settings. Example - {"model_group_retry_policy": {"max_retries": 5}}. IF null or {} then no router settings.
|
||||
|
||||
```
|
||||
curl --location 'http://0.0.0.0:4000/team/update' \
|
||||
|
|
@ -1396,6 +1401,10 @@ async def update_team( # noqa: PLR0915
|
|||
if _model_id is not None:
|
||||
updated_kv["model_id"] = _model_id
|
||||
|
||||
# Serialize router_settings to JSON if present (matching key update pattern)
|
||||
if "router_settings" in updated_kv and updated_kv["router_settings"] is not None:
|
||||
updated_kv["router_settings"] = safe_dumps(updated_kv["router_settings"])
|
||||
|
||||
updated_kv = prisma_client.jsonify_team_object(db_data=updated_kv)
|
||||
team_row: Optional[LiteLLM_TeamTable] = (
|
||||
await prisma_client.db.litellm_teamtable.update(
|
||||
|
|
|
|||
|
|
@ -124,6 +124,7 @@ model LiteLLM_TeamTable {
|
|||
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
|
||||
model_spend Json @default("{}")
|
||||
model_max_budget Json @default("{}")
|
||||
router_settings Json? @default("{}")
|
||||
team_member_permissions String[] @default([])
|
||||
model_id Int? @unique // id for LiteLLM_ModelTable -> stores team-level model aliases
|
||||
litellm_organization_table LiteLLM_OrganizationTable? @relation(fields: [organization_id], references: [organization_id])
|
||||
|
|
@ -225,6 +226,7 @@ model LiteLLM_VerificationToken {
|
|||
models String[]
|
||||
aliases Json @default("{}")
|
||||
config Json @default("{}")
|
||||
router_settings Json? @default("{}")
|
||||
user_id String?
|
||||
team_id String?
|
||||
permissions Json @default("{}")
|
||||
|
|
|
|||
|
|
@ -59,7 +59,7 @@ websockets = {version = "^15.0.1", optional = true}
|
|||
boto3 = {version = "1.36.0", optional = true}
|
||||
redisvl = {version = "^0.4.1", optional = true, markers = "python_version >= '3.9' and python_version < '3.14'"}
|
||||
mcp = {version = "^1.21.2", optional = true, python = ">=3.10"}
|
||||
litellm-proxy-extras = {version = "0.4.19", optional = true}
|
||||
litellm-proxy-extras = {version = "0.4.20", optional = true}
|
||||
rich = {version = "13.7.1", optional = true}
|
||||
litellm-enterprise = {version = "0.1.27", optional = true}
|
||||
diskcache = {version = "^5.6.1", optional = true}
|
||||
|
|
|
|||
|
|
@ -47,7 +47,7 @@ sentry_sdk==2.21.0 # for sentry error handling
|
|||
detect-secrets==1.5.0 # Enterprise - secret detection / masking in LLM requests
|
||||
cryptography==44.0.1
|
||||
tzdata==2025.1 # IANA time zone database
|
||||
litellm-proxy-extras==0.4.19 # for proxy extras - e.g. prisma migrations
|
||||
litellm-proxy-extras==0.4.20 # for proxy extras - e.g. prisma migrations
|
||||
llm-sandbox==0.3.31 # for skill execution in sandbox
|
||||
### LITELLM PACKAGE DEPENDENCIES
|
||||
python-dotenv==1.0.1 # for env
|
||||
|
|
|
|||
|
|
@ -124,6 +124,7 @@ model LiteLLM_TeamTable {
|
|||
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
|
||||
model_spend Json @default("{}")
|
||||
model_max_budget Json @default("{}")
|
||||
router_settings Json? @default("{}")
|
||||
team_member_permissions String[] @default([])
|
||||
model_id Int? @unique // id for LiteLLM_ModelTable -> stores team-level model aliases
|
||||
litellm_organization_table LiteLLM_OrganizationTable? @relation(fields: [organization_id], references: [organization_id])
|
||||
|
|
@ -225,6 +226,7 @@ model LiteLLM_VerificationToken {
|
|||
models String[]
|
||||
aliases Json @default("{}")
|
||||
config Json @default("{}")
|
||||
router_settings Json? @default("{}")
|
||||
user_id String?
|
||||
team_id String?
|
||||
permissions Json @default("{}")
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ import os
|
|||
import sys
|
||||
|
||||
import pytest
|
||||
import yaml
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
sys.path.insert(
|
||||
|
|
@ -3642,3 +3643,152 @@ async def test_update_key_negative_max_budget():
|
|||
# Should not raise any errors at model level
|
||||
request = UpdateKeyRequest(key="test-key", max_budget=-5.0)
|
||||
assert request.max_budget == -5.0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_generate_key_with_router_settings(monkeypatch):
|
||||
"""
|
||||
Test that /key/generate correctly handles router_settings by:
|
||||
1. Accepting router_settings as a dict parameter
|
||||
2. Serializing router_settings to JSON when saving to database
|
||||
3. Storing router_settings in the key record
|
||||
"""
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_prisma_client.jsonify_object = lambda data: data
|
||||
|
||||
# Mock prisma_client.insert_data for both user and key tables
|
||||
async def _insert_data_side_effect(*args, **kwargs):
|
||||
table_name = kwargs.get("table_name")
|
||||
if table_name == "user":
|
||||
return MagicMock(models=[], spend=0)
|
||||
elif table_name == "key":
|
||||
return MagicMock(
|
||||
token="hashed_token_router",
|
||||
litellm_budget_table=None,
|
||||
object_permission=None,
|
||||
)
|
||||
return MagicMock()
|
||||
|
||||
mock_prisma_client.insert_data = AsyncMock(side_effect=_insert_data_side_effect)
|
||||
mock_prisma_client.db = MagicMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken = MagicMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
|
||||
return_value=[]
|
||||
)
|
||||
mock_prisma_client.db.litellm_verificationtoken.count = AsyncMock(return_value=0)
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
|
||||
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 (
|
||||
generate_key_fn,
|
||||
)
|
||||
|
||||
# Test router_settings with sample data
|
||||
# Using valid UpdateRouterConfig fields (retry_policy is not a valid field,
|
||||
# but model_group_retry_policy is, which also tests nested dict serialization)
|
||||
router_settings_data = {
|
||||
"routing_strategy": "usage-based",
|
||||
"num_retries": 3,
|
||||
"model_group_retry_policy": {"max_retries": 5},
|
||||
}
|
||||
|
||||
request_data = GenerateKeyRequest(
|
||||
models=["gpt-4"],
|
||||
router_settings=router_settings_data,
|
||||
)
|
||||
|
||||
await generate_key_fn(
|
||||
data=request_data,
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-1234",
|
||||
user_id="user-router-1",
|
||||
),
|
||||
)
|
||||
|
||||
# Verify key insertion was called
|
||||
assert mock_prisma_client.insert_data.call_count >= 1
|
||||
key_insert_calls = [
|
||||
call.kwargs
|
||||
for call in mock_prisma_client.insert_data.call_args_list
|
||||
if call.kwargs.get("table_name") == "key"
|
||||
]
|
||||
assert len(key_insert_calls) >= 1
|
||||
key_data = key_insert_calls[0]["data"]
|
||||
|
||||
# Verify router_settings is present
|
||||
assert "router_settings" in key_data
|
||||
|
||||
# router_settings should be present in the data passed to insert_data
|
||||
# The code uses safe_dumps to serialize router_settings, so it will be a JSON string
|
||||
router_settings_value = key_data["router_settings"]
|
||||
|
||||
# Get the actual settings value for comparison
|
||||
# The code uses safe_dumps to serialize and yaml.safe_load to deserialize
|
||||
if isinstance(router_settings_value, str):
|
||||
# If it's a JSON string (from safe_dumps), deserialize it using json.loads
|
||||
# (safe_dumps produces JSON, and json.loads is the correct way to deserialize it)
|
||||
actual_settings = json.loads(router_settings_value)
|
||||
elif isinstance(router_settings_value, dict):
|
||||
# If it's still a dict, use it directly
|
||||
actual_settings = router_settings_value
|
||||
else:
|
||||
raise AssertionError(
|
||||
f"router_settings should be str or dict, got {type(router_settings_value)}"
|
||||
)
|
||||
|
||||
# Verify router_settings matches input (regardless of serialization state)
|
||||
assert actual_settings == router_settings_data
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_key_with_router_settings(monkeypatch):
|
||||
"""
|
||||
Test that /key/update correctly handles router_settings by:
|
||||
1. Accepting router_settings as a dict parameter
|
||||
2. Serializing router_settings to JSON when updating database
|
||||
3. Updating router_settings in the key record
|
||||
"""
|
||||
from litellm.proxy._types import LiteLLM_VerificationToken, UpdateKeyRequest
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
prepare_key_update_data,
|
||||
)
|
||||
|
||||
# Mock existing key
|
||||
existing_key = LiteLLM_VerificationToken(
|
||||
token="test-token-router",
|
||||
key_alias="test-key",
|
||||
models=["gpt-3.5-turbo"],
|
||||
user_id="test-user",
|
||||
team_id=None,
|
||||
auto_rotate=False,
|
||||
rotation_interval=None,
|
||||
metadata={},
|
||||
)
|
||||
|
||||
# Test updating router_settings
|
||||
router_settings_data = {
|
||||
"routing_strategy": "latency-based",
|
||||
"num_retries": 2,
|
||||
}
|
||||
|
||||
update_request = UpdateKeyRequest(
|
||||
key="test-token-router", router_settings=router_settings_data
|
||||
)
|
||||
|
||||
result = await prepare_key_update_data(
|
||||
data=update_request, existing_key_row=existing_key
|
||||
)
|
||||
|
||||
# Verify router_settings is serialized to JSON string
|
||||
assert "router_settings" in result
|
||||
assert isinstance(result["router_settings"], str)
|
||||
|
||||
# Verify router_settings can be deserialized and matches input
|
||||
deserialized_settings = json.loads(result["router_settings"])
|
||||
assert deserialized_settings == router_settings_data
|
||||
|
|
|
|||
|
|
@ -4393,3 +4393,162 @@ async def test_new_team_positive_budgets_accepted():
|
|||
)
|
||||
assert request.max_budget == 100.0
|
||||
assert request.team_member_budget == 50.0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_new_team_with_router_settings(mock_db_client, mock_admin_auth):
|
||||
"""
|
||||
Test that /team/new correctly handles router_settings by:
|
||||
1. Accepting router_settings as a dict parameter
|
||||
2. Serializing router_settings to JSON when saving to database
|
||||
3. Storing router_settings in the team record
|
||||
"""
|
||||
# Configure mocked prisma client
|
||||
mock_db_client.jsonify_team_object = lambda db_data: db_data
|
||||
mock_db_client.get_data = AsyncMock(return_value=None)
|
||||
mock_db_client.update_data = AsyncMock(return_value=MagicMock())
|
||||
mock_db_client.db = MagicMock()
|
||||
|
||||
# Mock model table creation
|
||||
mock_db_client.db.litellm_modeltable = MagicMock()
|
||||
mock_db_client.db.litellm_modeltable.create = AsyncMock(
|
||||
return_value=MagicMock(id="model123")
|
||||
)
|
||||
|
||||
# Capture team table creation
|
||||
team_create_result = MagicMock(
|
||||
team_id="team-router-456",
|
||||
)
|
||||
team_create_result.model_dump.return_value = {
|
||||
"team_id": "team-router-456",
|
||||
}
|
||||
mock_team_create = AsyncMock(return_value=team_create_result)
|
||||
mock_team_count = AsyncMock(return_value=0)
|
||||
mock_db_client.db.litellm_teamtable = MagicMock()
|
||||
mock_db_client.db.litellm_teamtable.create = mock_team_create
|
||||
mock_db_client.db.litellm_teamtable.count = mock_team_count
|
||||
mock_db_client.db.litellm_teamtable.update = AsyncMock(
|
||||
return_value=team_create_result
|
||||
)
|
||||
|
||||
# Mock user table
|
||||
mock_db_client.db.litellm_usertable = MagicMock()
|
||||
mock_db_client.db.litellm_usertable.update = AsyncMock(return_value=MagicMock())
|
||||
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._types import NewTeamRequest
|
||||
from litellm.proxy.management_endpoints.team_endpoints import new_team
|
||||
|
||||
# Test router_settings with sample data
|
||||
router_settings_data = {
|
||||
"routing_strategy": "usage-based",
|
||||
"num_retries": 3,
|
||||
"retry_policy": {"max_retries": 5},
|
||||
}
|
||||
|
||||
# Build request with router_settings
|
||||
team_request = NewTeamRequest(
|
||||
team_alias="my-team-router",
|
||||
router_settings=router_settings_data,
|
||||
)
|
||||
|
||||
dummy_request = MagicMock(spec=Request)
|
||||
|
||||
# Execute the endpoint function
|
||||
await new_team(
|
||||
data=team_request,
|
||||
http_request=dummy_request,
|
||||
user_api_key_dict=mock_admin_auth,
|
||||
)
|
||||
|
||||
# Verify team creation was called
|
||||
assert mock_team_create.call_count == 1
|
||||
created_team_kwargs = mock_team_create.call_args.kwargs
|
||||
team_data = created_team_kwargs["data"]
|
||||
|
||||
# Verify router_settings is serialized to JSON string
|
||||
assert "router_settings" in team_data
|
||||
assert isinstance(team_data["router_settings"], str)
|
||||
|
||||
# Verify router_settings can be deserialized and matches input
|
||||
deserialized_settings = json.loads(team_data["router_settings"])
|
||||
assert deserialized_settings == router_settings_data
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_team_with_router_settings(mock_db_client, mock_admin_auth):
|
||||
"""
|
||||
Test that /team/update correctly handles router_settings by:
|
||||
1. Accepting router_settings as a dict parameter
|
||||
2. Serializing router_settings to JSON when updating database
|
||||
3. Updating router_settings in the team record
|
||||
"""
|
||||
# Configure mocked prisma client
|
||||
mock_db_client.jsonify_team_object = lambda db_data: db_data
|
||||
mock_db_client.db = MagicMock()
|
||||
|
||||
# Mock existing team row
|
||||
existing_team_mock = MagicMock()
|
||||
existing_team_mock.team_id = "team-router-update-789"
|
||||
existing_team_mock.organization_id = None
|
||||
existing_team_mock.models = []
|
||||
existing_team_mock.members_with_roles = []
|
||||
existing_team_mock.model_dump.return_value = {
|
||||
"team_id": "team-router-update-789",
|
||||
"organization_id": None,
|
||||
"models": [],
|
||||
"members_with_roles": [],
|
||||
}
|
||||
|
||||
# Mock team table find_unique and update
|
||||
updated_team_result = MagicMock(
|
||||
team_id="team-router-update-789",
|
||||
)
|
||||
updated_team_result.model_dump.return_value = {
|
||||
"team_id": "team-router-update-789",
|
||||
}
|
||||
mock_team_find_unique = AsyncMock(return_value=existing_team_mock)
|
||||
mock_team_update = AsyncMock(return_value=updated_team_result)
|
||||
mock_db_client.db.litellm_teamtable = MagicMock()
|
||||
mock_db_client.db.litellm_teamtable.find_unique = mock_team_find_unique
|
||||
mock_db_client.db.litellm_teamtable.update = mock_team_update
|
||||
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._types import UpdateTeamRequest
|
||||
from litellm.proxy.management_endpoints.team_endpoints import update_team
|
||||
|
||||
# Test router_settings with updated data
|
||||
router_settings_data = {
|
||||
"routing_strategy": "latency-based",
|
||||
"num_retries": 2,
|
||||
}
|
||||
|
||||
# Build update request with router_settings
|
||||
team_update_request = UpdateTeamRequest(
|
||||
team_id="team-router-update-789",
|
||||
router_settings=router_settings_data,
|
||||
)
|
||||
|
||||
dummy_request = MagicMock(spec=Request)
|
||||
|
||||
# Execute the endpoint function
|
||||
await update_team(
|
||||
data=team_update_request,
|
||||
http_request=dummy_request,
|
||||
user_api_key_dict=mock_admin_auth,
|
||||
)
|
||||
|
||||
# Verify team update was called
|
||||
assert mock_team_update.call_count == 1
|
||||
updated_team_kwargs = mock_team_update.call_args.kwargs
|
||||
team_data = updated_team_kwargs["data"]
|
||||
|
||||
# Verify router_settings is serialized to JSON string
|
||||
assert "router_settings" in team_data
|
||||
assert isinstance(team_data["router_settings"], str)
|
||||
|
||||
# Verify router_settings can be deserialized and matches input
|
||||
deserialized_settings = json.loads(team_data["router_settings"])
|
||||
assert deserialized_settings == router_settings_data
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue