From ca97ea8acdb74c84b72bfb689a1909bc1670fcf4 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Wed, 6 Mar 2024 17:42:08 -0800 Subject: [PATCH 1/4] feat(proxy_server.py): team based model aliases allow setting model aliases at a team level (e.g. route all 'gpt-3.5-turbo' requests from team-1 to model-deployment-group-2) --- litellm/proxy/_types.py | 17 +++++++++++++++-- litellm/proxy/proxy_server.py | 28 ++++++++++++++++++++++++++-- litellm/proxy/schema.prisma | 13 +++++++++++++ litellm/proxy/utils.py | 15 ++++++++++++--- schema.prisma | 13 +++++++++++++ 5 files changed, 79 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 7ae67bdc638..fd85280ddd1 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -212,6 +212,12 @@ class KeyRequest(LiteLLMBase): keys: List[str] +class LiteLLM_ModelTable(LiteLLMBase): + model_aliases: Optional[str] = None # json dump the dict + created_by: str + updated_by: str + + class NewUserRequest(GenerateKeyRequest): max_budget: Optional[float] = None user_email: Optional[str] = None @@ -251,7 +257,7 @@ class Member(LiteLLMBase): return values -class NewTeamRequest(LiteLLMBase): +class TeamBase(LiteLLMBase): team_alias: Optional[str] = None team_id: Optional[str] = None organization_id: Optional[str] = None @@ -265,6 +271,10 @@ class NewTeamRequest(LiteLLMBase): models: list = [] +class NewTeamRequest(TeamBase): + model_aliases: Optional[dict] = None + + class GlobalEndUsersSpend(LiteLLMBase): api_key: Optional[str] = None @@ -299,11 +309,12 @@ class DeleteTeamRequest(LiteLLMBase): team_ids: List[str] # required -class LiteLLM_TeamTable(NewTeamRequest): +class LiteLLM_TeamTable(TeamBase): spend: Optional[float] = None max_parallel_requests: Optional[int] = None budget_duration: Optional[str] = None budget_reset_at: Optional[datetime] = None + model_id: Optional[int] = None @root_validator(pre=True) def set_model_info(cls, values): @@ -313,6 +324,7 @@ class LiteLLM_TeamTable(NewTeamRequest): "config", "permissions", "model_max_budget", + "model_aliases", ] for field in dict_fields: value = values.get(field) @@ -542,6 +554,7 @@ class LiteLLM_VerificationTokenView(LiteLLM_VerificationToken): team_rpm_limit: Optional[int] = None team_max_budget: Optional[float] = None soft_budget: Optional[float] = None + team_model_aliases: Optional[Dict] = None class UserAPIKeyAuth( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 409cf63d550..5ee3b751ff6 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -405,7 +405,15 @@ async def user_api_key_auth( ) # request data, used across all checks. Making this easily available # Check 1. If token can call model - litellm.model_alias_map = valid_token.aliases + _model_alias_map = {} + if valid_token.team_model_aliases is not None: + _model_alias_map = { + **valid_token.aliases, + **valid_token.team_model_aliases, + } + else: + _model_alias_map = {**valid_token.aliases} + litellm.model_alias_map = _model_alias_map config = valid_token.config if config != {}: model_list = config.get("model_list", []) @@ -5020,11 +5028,27 @@ async def new_team( Member(role="admin", user_id=user_api_key_dict.user_id) ) + ## ADD TO MODEL TABLE + _model_id = None + if data.model_aliases is not None and isinstance(data.model_aliases, dict): + litellm_modeltable = LiteLLM_ModelTable( + model_aliases=json.dumps(data.model_aliases), + created_by=user_api_key_dict.user_id or litellm_proxy_admin_name, + updated_by=user_api_key_dict.user_id or litellm_proxy_admin_name, + ) + model_dict = await prisma_client.db.litellm_modeltable.create( + {**litellm_modeltable.json(exclude_none=True)} # type: ignore + ) # type: ignore + + _model_id = model_dict.id + + ## ADD TO TEAM TABLE complete_team_data = LiteLLM_TeamTable( **data.json(), max_parallel_requests=user_api_key_dict.max_parallel_requests, budget_duration=user_api_key_dict.budget_duration, budget_reset_at=user_api_key_dict.budget_reset_at, + model_id=_model_id, ) team_row = await prisma_client.insert_data( @@ -5495,7 +5519,7 @@ async def new_organization( - `organization_alias`: *str* = The name of the organization. - `models`: *List* = The models the organization has access to. - `budget_id`: *Optional[str]* = The id for a budget (tpm/rpm/max budget) for the organization. - ### IF NO BUDGET - CREATE ONE WITH THESE PARAMS ### + ### IF NO BUDGET ID - CREATE ONE WITH THESE PARAMS ### - `max_budget`: *Optional[float]* = Max budget for org - `tpm_limit`: *Optional[int]* = Max tpm limit for org - `rpm_limit`: *Optional[int]* = Max rpm limit for org diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 265bf32c076..d8c8faf1606 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -42,6 +42,17 @@ model LiteLLM_OrganizationTable { teams LiteLLM_TeamTable[] } +// Model info for teams, just has model aliases for now. +model LiteLLM_ModelTable { + id Int @id @default(autoincrement()) + model_aliases Json? @map("aliases") + created_at DateTime @default(now()) @map("created_at") + created_by String + updated_at DateTime @default(now()) @updatedAt @map("updated_at") + updated_by String + team LiteLLM_TeamTable? +} + // Assign prod keys to groups, not individuals model LiteLLM_TeamTable { team_id String @id @default(uuid()) @@ -63,7 +74,9 @@ model LiteLLM_TeamTable { updated_at DateTime @default(now()) @updatedAt @map("updated_at") model_spend Json @default("{}") model_max_budget Json @default("{}") + model_id Int? @unique litellm_organization_table LiteLLM_OrganizationTable? @relation(fields: [organization_id], references: [organization_id]) + litellm_model_table LiteLLM_ModelTable? @relation(fields: [model_id], references: [id]) } // Track spend, rate limit, budget Users diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 1e701515e1b..ee5e323e8ef 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -965,12 +965,21 @@ class PrismaClient: ) sql_query = f""" - SELECT * - FROM "LiteLLM_VerificationTokenView" - WHERE token = '{token}' + SELECT + v.*, + t.spend AS team_spend, + t.max_budget AS team_max_budget, + t.tpm_limit AS team_tpm_limit, + t.rpm_limit AS team_rpm_limit, + m.aliases as team_model_aliases + FROM "LiteLLM_VerificationToken" AS v + INNER JOIN "LiteLLM_TeamTable" AS t ON v.team_id = t.team_id + LEFT JOIN "LiteLLM_ModelTable" m ON t.model_id = m.id + WHERE v.token = '{token}' """ response = await self.db.query_first(query=sql_query) + if response is not None: response = LiteLLM_VerificationTokenView(**response) # for prisma we need to cast the expires time to str diff --git a/schema.prisma b/schema.prisma index 265bf32c076..d8c8faf1606 100644 --- a/schema.prisma +++ b/schema.prisma @@ -42,6 +42,17 @@ model LiteLLM_OrganizationTable { teams LiteLLM_TeamTable[] } +// Model info for teams, just has model aliases for now. +model LiteLLM_ModelTable { + id Int @id @default(autoincrement()) + model_aliases Json? @map("aliases") + created_at DateTime @default(now()) @map("created_at") + created_by String + updated_at DateTime @default(now()) @updatedAt @map("updated_at") + updated_by String + team LiteLLM_TeamTable? +} + // Assign prod keys to groups, not individuals model LiteLLM_TeamTable { team_id String @id @default(uuid()) @@ -63,7 +74,9 @@ model LiteLLM_TeamTable { updated_at DateTime @default(now()) @updatedAt @map("updated_at") model_spend Json @default("{}") model_max_budget Json @default("{}") + model_id Int? @unique litellm_organization_table LiteLLM_OrganizationTable? @relation(fields: [organization_id], references: [organization_id]) + litellm_model_table LiteLLM_ModelTable? @relation(fields: [model_id], references: [id]) } // Track spend, rate limit, budget Users From d1d8adfb115095701bbd158a9224a4e499f0ad45 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Wed, 6 Mar 2024 19:41:12 -0800 Subject: [PATCH 2/4] fix(proxy_server.py): fix sql query --- litellm/proxy/proxy_server.py | 5 ++++- litellm/proxy/utils.py | 2 +- 2 files changed, 5 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 5ee3b751ff6..bb4bfe47a05 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -831,7 +831,10 @@ async def user_api_key_auth( raise Exception( f"This key is made for LiteLLM UI, Tried to access route: {route}. Not allowed" ) - return UserAPIKeyAuth(api_key=api_key, **valid_token_dict) + if valid_token_dict is not None: + return UserAPIKeyAuth(api_key=api_key, **valid_token_dict) + else: + raise Exception() except Exception as e: # verbose_proxy_logger.debug(f"An exception occurred - {traceback.format_exc()}") traceback.print_exc() diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index ee5e323e8ef..81130787d07 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -973,7 +973,7 @@ class PrismaClient: t.rpm_limit AS team_rpm_limit, m.aliases as team_model_aliases FROM "LiteLLM_VerificationToken" AS v - INNER JOIN "LiteLLM_TeamTable" AS t ON v.team_id = t.team_id + LEFT JOIN "LiteLLM_TeamTable" AS t ON v.team_id = t.team_id LEFT JOIN "LiteLLM_ModelTable" m ON t.model_id = m.id WHERE v.token = '{token}' """ From be6674f5e0715102aba56168508d67bdfe5ea054 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Wed, 6 Mar 2024 20:13:09 -0800 Subject: [PATCH 3/4] test(test_completion.py): fix test --- litellm/tests/test_completion.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/litellm/tests/test_completion.py b/litellm/tests/test_completion.py index af00275d3a1..0643a8bef40 100644 --- a/litellm/tests/test_completion.py +++ b/litellm/tests/test_completion.py @@ -1453,9 +1453,9 @@ def test_completion_replicate_vicuna(): def test_replicate_custom_prompt_dict(): litellm.set_verbose = True - model_name = "replicate/meta/llama-2-7b-chat:13c3cdee13ee059ab779f0291d29054dab00a47dad8261375654de5540165fb0" + model_name = "replicate/meta/llama-2-7b-chat" litellm.register_prompt_template( - model="replicate/meta/llama-2-7b-chat:13c3cdee13ee059ab779f0291d29054dab00a47dad8261375654de5540165fb0", + model="replicate/meta/llama-2-7b-chat", initial_prompt_value="You are a good assistant", # [OPTIONAL] roles={ "system": { From c0c3117dec426696a98d8c93ae0614c9629de7fc Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Wed, 6 Mar 2024 20:47:05 -0800 Subject: [PATCH 4/4] fix(utils.py): fix get optional param embeddings --- litellm/utils.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/litellm/utils.py b/litellm/utils.py index 68dc137afbc..a1f1bb374c5 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -4030,11 +4030,11 @@ def get_optional_params_embeddings( keys = list(non_default_params.keys()) for k in keys: non_default_params.pop(k, None) - return non_default_params - raise UnsupportedParamsError( - status_code=500, - message=f"Setting user/encoding format is not supported by {custom_llm_provider}. To drop it from the call, set `litellm.drop_params = True`.", - ) + else: + raise UnsupportedParamsError( + status_code=500, + message=f"Setting user/encoding format is not supported by {custom_llm_provider}. To drop it from the call, set `litellm.drop_params = True`.", + ) final_params = {**non_default_params, **kwargs} return final_params