mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-09 00:26:01 +00:00
[Fix] QA Fixes - Vector Store Object Permissions (#11291)
* fix: QA for key,team,org permissions * fix: add_vector_store_to_registry * fix: refactor bedrock guard * fix: refactor using us east 1 with vector stores * fix: code QA checks * fix: testing for mgmt endpoints
This commit is contained in:
@@ -34,7 +34,6 @@ from litellm.types.vector_stores import (
|
||||
VectorStoreSearchResponse,
|
||||
VectorStoreSearchResult,
|
||||
)
|
||||
from litellm.utils import load_credentials_from_list
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
@@ -258,22 +257,49 @@ class BedrockVectorStore(BaseVectorStore, BaseAWSLLM):
|
||||
from fastapi import HTTPException
|
||||
|
||||
non_default_params = non_default_params or {}
|
||||
load_credentials_from_list(kwargs=non_default_params)
|
||||
credentials_dict: Dict[str, Any] = {}
|
||||
if litellm.vector_store_registry is not None:
|
||||
credentials_dict = (
|
||||
litellm.vector_store_registry.get_credentials_for_vector_store(
|
||||
knowledge_base_id
|
||||
)
|
||||
)
|
||||
|
||||
credentials = self.get_credentials(
|
||||
aws_access_key_id=non_default_params.get("aws_access_key_id", None),
|
||||
aws_secret_access_key=non_default_params.get("aws_secret_access_key", None),
|
||||
aws_session_token=non_default_params.get("aws_session_token", None),
|
||||
aws_region_name=non_default_params.get("aws_region_name", None),
|
||||
aws_session_name=non_default_params.get("aws_session_name", None),
|
||||
aws_profile_name=non_default_params.get("aws_profile_name", None),
|
||||
aws_role_name=non_default_params.get("aws_role_name", None),
|
||||
aws_web_identity_token=non_default_params.get(
|
||||
"aws_web_identity_token", None
|
||||
aws_access_key_id=credentials_dict.get(
|
||||
"aws_access_key_id", non_default_params.get("aws_access_key_id", None)
|
||||
),
|
||||
aws_secret_access_key=credentials_dict.get(
|
||||
"aws_secret_access_key",
|
||||
non_default_params.get("aws_secret_access_key", None),
|
||||
),
|
||||
aws_session_token=credentials_dict.get(
|
||||
"aws_session_token", non_default_params.get("aws_session_token", None)
|
||||
),
|
||||
aws_region_name=credentials_dict.get(
|
||||
"aws_region_name", non_default_params.get("aws_region_name", None)
|
||||
),
|
||||
aws_session_name=credentials_dict.get(
|
||||
"aws_session_name", non_default_params.get("aws_session_name", None)
|
||||
),
|
||||
aws_profile_name=credentials_dict.get(
|
||||
"aws_profile_name", non_default_params.get("aws_profile_name", None)
|
||||
),
|
||||
aws_role_name=credentials_dict.get(
|
||||
"aws_role_name", non_default_params.get("aws_role_name", None)
|
||||
),
|
||||
aws_web_identity_token=credentials_dict.get(
|
||||
"aws_web_identity_token",
|
||||
non_default_params.get("aws_web_identity_token", None),
|
||||
),
|
||||
aws_sts_endpoint=credentials_dict.get(
|
||||
"aws_sts_endpoint", non_default_params.get("aws_sts_endpoint", None)
|
||||
),
|
||||
aws_sts_endpoint=non_default_params.get("aws_sts_endpoint", None),
|
||||
)
|
||||
aws_region_name = self._get_aws_region_name(
|
||||
optional_params=self.optional_params
|
||||
aws_region_name = self.get_aws_region_name_for_non_llm_api_calls(
|
||||
aws_region_name=credentials_dict.get(
|
||||
"aws_region_name", non_default_params.get("aws_region_name", None)
|
||||
),
|
||||
)
|
||||
|
||||
# Prepare request data
|
||||
|
||||
@@ -336,6 +336,36 @@ class BaseAWSLLM:
|
||||
|
||||
return aws_region_name
|
||||
|
||||
def get_aws_region_name_for_non_llm_api_calls(
|
||||
self,
|
||||
aws_region_name: Optional[str] = None,
|
||||
):
|
||||
"""
|
||||
Get the AWS region name for non-llm api calls.
|
||||
|
||||
LLM API calls check the model arn and end up using that as the region name.
|
||||
|
||||
For non-llm api calls eg. Guardrails, Vector Stores we just need to check the dynamic param or env vars.
|
||||
"""
|
||||
if aws_region_name is None:
|
||||
# check env #
|
||||
litellm_aws_region_name = get_secret("AWS_REGION_NAME", None)
|
||||
|
||||
if litellm_aws_region_name is not None and isinstance(
|
||||
litellm_aws_region_name, str
|
||||
):
|
||||
aws_region_name = litellm_aws_region_name
|
||||
|
||||
standard_aws_region_name = get_secret("AWS_REGION", None)
|
||||
if standard_aws_region_name is not None and isinstance(
|
||||
standard_aws_region_name, str
|
||||
):
|
||||
aws_region_name = standard_aws_region_name
|
||||
|
||||
if aws_region_name is None:
|
||||
aws_region_name = "us-west-2"
|
||||
return aws_region_name
|
||||
|
||||
@tracer.wrap()
|
||||
def _auth_with_web_identity_token(
|
||||
self,
|
||||
|
||||
@@ -30,7 +30,6 @@ from litellm.llms.custom_httpx.http_handler import (
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.secret_managers.main import get_secret
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
|
||||
@@ -129,23 +128,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
||||
aws_sts_endpoint = self.optional_params.get("aws_sts_endpoint", None)
|
||||
|
||||
### SET REGION NAME ###
|
||||
if aws_region_name is None:
|
||||
# check env #
|
||||
litellm_aws_region_name = get_secret("AWS_REGION_NAME", None)
|
||||
|
||||
if litellm_aws_region_name is not None and isinstance(
|
||||
litellm_aws_region_name, str
|
||||
):
|
||||
aws_region_name = litellm_aws_region_name
|
||||
|
||||
standard_aws_region_name = get_secret("AWS_REGION", None)
|
||||
if standard_aws_region_name is not None and isinstance(
|
||||
standard_aws_region_name, str
|
||||
):
|
||||
aws_region_name = standard_aws_region_name
|
||||
|
||||
if aws_region_name is None:
|
||||
aws_region_name = "us-west-2"
|
||||
aws_region_name = self.get_aws_region_name_for_non_llm_api_calls(
|
||||
aws_region_name=aws_region_name,
|
||||
)
|
||||
|
||||
credentials: Credentials = self.get_credentials(
|
||||
aws_access_key_id=aws_access_key_id,
|
||||
|
||||
@@ -45,6 +45,9 @@ from litellm.proxy.management_endpoints.common_utils import (
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
_add_model_to_db,
|
||||
)
|
||||
from litellm.proxy.management_helpers.object_permission_utils import (
|
||||
handle_update_object_permission_common,
|
||||
)
|
||||
from litellm.proxy.management_helpers.team_member_permission_checks import (
|
||||
TeamMemberPermissionChecks,
|
||||
)
|
||||
@@ -733,48 +736,20 @@ async def _handle_update_object_permission(
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise ValueError("Prisma client not found")
|
||||
|
||||
#########################################################
|
||||
# Ensure `object_permission` is not added to the data_json
|
||||
# We need to update the entity at the object_permission_id level in the LiteLLM_ObjectPermissionTable
|
||||
#########################################################
|
||||
new_object_permission = data_json.pop("object_permission")
|
||||
if new_object_permission is None:
|
||||
return data_json
|
||||
|
||||
# lookup existing object permission ID and update that entry
|
||||
existing_object_permission_id = existing_key_row.object_permission_id
|
||||
existing_object_permission = (
|
||||
await prisma_client.db.litellm_objectpermissiontable.find_unique(
|
||||
where={"object_permission_id": existing_object_permission_id},
|
||||
)
|
||||
# Use the common helper to handle the object permission update
|
||||
object_permission_id = await handle_update_object_permission_common(
|
||||
data_json=data_json,
|
||||
existing_object_permission_id=existing_key_row.object_permission_id,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
existing_object_permissions_dict: dict = {}
|
||||
if existing_object_permission is not None:
|
||||
# update the object permission
|
||||
existing_object_permissions_dict = existing_object_permission.model_dump(
|
||||
exclude_unset=True, exclude_none=True
|
||||
)
|
||||
existing_object_permissions_dict.update(dict(new_object_permission))
|
||||
|
||||
#########################################################
|
||||
# Commit the update to the LiteLLM_ObjectPermissionTable
|
||||
#########################################################
|
||||
new_object_permission_row = (
|
||||
await prisma_client.db.litellm_objectpermissiontable.upsert(
|
||||
where={"object_permission_id": existing_object_permission_id},
|
||||
data={
|
||||
"create": existing_object_permissions_dict,
|
||||
"update": existing_object_permissions_dict,
|
||||
},
|
||||
# Add the object_permission_id to data_json if one was created/updated
|
||||
if object_permission_id is not None:
|
||||
data_json["object_permission_id"] = object_permission_id
|
||||
verbose_proxy_logger.debug(
|
||||
f"updated object_permission_id: {object_permission_id}"
|
||||
)
|
||||
)
|
||||
verbose_proxy_logger.debug(
|
||||
f"new_object_permission_row: {new_object_permission_row}"
|
||||
)
|
||||
data_json["object_permission_id"] = new_object_permission_row.object_permission_id
|
||||
|
||||
return data_json
|
||||
|
||||
|
||||
|
||||
@@ -24,6 +24,9 @@ from litellm.proxy.management_endpoints.budget_management_endpoints import (
|
||||
new_budget,
|
||||
update_budget,
|
||||
)
|
||||
from litellm.proxy.management_helpers.object_permission_utils import (
|
||||
handle_update_object_permission_common,
|
||||
)
|
||||
from litellm.proxy.management_helpers.utils import (
|
||||
get_new_internal_user_defaults,
|
||||
management_endpoint_wrapper,
|
||||
@@ -309,61 +312,20 @@ async def handle_update_object_permission(
|
||||
- Upserts the new object permission into the LiteLLM_ObjectPermissionTable
|
||||
- Adds object_permission_id to data_json (this gets added in the DB)
|
||||
- Pops the object_permission from data_json
|
||||
-
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise ValueError("Prisma client not found")
|
||||
|
||||
#########################################################
|
||||
# Ensure `object_permission` is not added to the data_json
|
||||
# We need to update the entity at the object_permission_id level in the LiteLLM_ObjectPermissionTable
|
||||
#########################################################
|
||||
new_object_permission: Union[dict, str] = data_json.pop("object_permission") or {}
|
||||
if new_object_permission is None:
|
||||
return data_json
|
||||
|
||||
# lookup existing object permission ID and update that entry
|
||||
existing_object_permission_id = existing_organization_row.object_permission_id
|
||||
existing_object_permissions_dict = {}
|
||||
|
||||
existing_object_permission = (
|
||||
await prisma_client.db.litellm_objectpermissiontable.find_unique(
|
||||
where={"object_permission_id": existing_object_permission_id},
|
||||
)
|
||||
# Use the common helper to handle the object permission update
|
||||
object_permission_id = await handle_update_object_permission_common(
|
||||
data_json=data_json,
|
||||
existing_object_permission_id=existing_organization_row.object_permission_id,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
# update the object permission
|
||||
if existing_object_permission is not None:
|
||||
existing_object_permissions_dict = existing_object_permission.model_dump(
|
||||
exclude_unset=True, exclude_none=True
|
||||
)
|
||||
# Add the object_permission_id to data_json if one was created/updated
|
||||
if object_permission_id is not None:
|
||||
data_json["object_permission_id"] = object_permission_id
|
||||
|
||||
if isinstance(new_object_permission, str):
|
||||
new_object_permission = json.loads(new_object_permission)
|
||||
|
||||
if isinstance(new_object_permission, dict):
|
||||
existing_object_permissions_dict.update(new_object_permission)
|
||||
|
||||
#########################################################
|
||||
# Commit the update to the LiteLLM_ObjectPermissionTable
|
||||
#########################################################
|
||||
created_object_permission_row = (
|
||||
await prisma_client.db.litellm_objectpermissiontable.upsert(
|
||||
where={"object_permission_id": existing_object_permission_id},
|
||||
data={
|
||||
"create": existing_object_permissions_dict,
|
||||
"update": existing_object_permissions_dict,
|
||||
},
|
||||
)
|
||||
)
|
||||
data_json[
|
||||
"object_permission_id"
|
||||
] = created_object_permission_row.object_permission_id
|
||||
verbose_proxy_logger.debug(
|
||||
f"created_object_permission_row: {created_object_permission_row}"
|
||||
)
|
||||
return data_json
|
||||
|
||||
|
||||
|
||||
@@ -70,6 +70,9 @@ from litellm.proxy.management_endpoints.common_utils import (
|
||||
from litellm.proxy.management_endpoints.tag_management_endpoints import (
|
||||
get_daily_activity,
|
||||
)
|
||||
from litellm.proxy.management_helpers.object_permission_utils import (
|
||||
handle_update_object_permission_common,
|
||||
)
|
||||
from litellm.proxy.management_helpers.team_member_permission_checks import (
|
||||
TeamMemberPermissionChecks,
|
||||
)
|
||||
@@ -752,48 +755,20 @@ async def handle_update_object_permission(
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise ValueError("Prisma client not found")
|
||||
|
||||
#########################################################
|
||||
# Ensure `object_permission` is not added to the data_json
|
||||
# We need to update the entity at the object_permission_id level in the LiteLLM_ObjectPermissionTable
|
||||
#########################################################
|
||||
new_object_permission = data_json.pop("object_permission")
|
||||
if new_object_permission is None:
|
||||
return data_json
|
||||
|
||||
# lookup existing object permission ID and update that entry
|
||||
existing_object_permission_id = existing_team_row.object_permission_id
|
||||
existing_object_permission = (
|
||||
await prisma_client.db.litellm_objectpermissiontable.find_unique(
|
||||
where={"object_permission_id": existing_object_permission_id},
|
||||
)
|
||||
# Use the common helper to handle the object permission update
|
||||
object_permission_id = await handle_update_object_permission_common(
|
||||
data_json=data_json,
|
||||
existing_object_permission_id=existing_team_row.object_permission_id,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
existing_object_permissions_dict: Dict = {}
|
||||
# Add the object_permission_id to data_json if one was created/updated
|
||||
if object_permission_id is not None:
|
||||
data_json["object_permission_id"] = object_permission_id
|
||||
verbose_proxy_logger.debug(
|
||||
f"updated object_permission_id: {object_permission_id}"
|
||||
)
|
||||
|
||||
# update the object permission
|
||||
if existing_object_permission is not None:
|
||||
existing_object_permissions_dict = existing_object_permission.model_dump(
|
||||
exclude_unset=True, exclude_none=True
|
||||
)
|
||||
existing_object_permissions_dict.update(dict(new_object_permission))
|
||||
created_object_permission_row = (
|
||||
await prisma_client.db.litellm_objectpermissiontable.upsert(
|
||||
where={"object_permission_id": existing_object_permission_id},
|
||||
data={
|
||||
"create": existing_object_permissions_dict,
|
||||
"update": existing_object_permissions_dict,
|
||||
},
|
||||
)
|
||||
)
|
||||
data_json[
|
||||
"object_permission_id"
|
||||
] = created_object_permission_row.object_permission_id
|
||||
verbose_proxy_logger.debug(
|
||||
f"created_object_permission_row: {created_object_permission_row}"
|
||||
)
|
||||
return data_json
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,93 @@
|
||||
"""
|
||||
Common utility functions for handling object permission updates across
|
||||
organizations, teams, and keys.
|
||||
"""
|
||||
|
||||
import json
|
||||
import uuid
|
||||
from typing import Dict, Optional, Union
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
|
||||
async def handle_update_object_permission_common(
|
||||
data_json: Dict,
|
||||
existing_object_permission_id: Optional[str],
|
||||
prisma_client: Optional[PrismaClient],
|
||||
) -> Optional[str]:
|
||||
"""
|
||||
Common logic for handling object permission updates across organizations, teams, and keys.
|
||||
|
||||
This function:
|
||||
1. Extracts `object_permission` from data_json
|
||||
2. Looks up existing object permission if it exists
|
||||
3. Merges new permissions with existing ones
|
||||
4. Upserts to the LiteLLM_ObjectPermissionTable
|
||||
5. Returns the object_permission_id
|
||||
|
||||
Args:
|
||||
data_json: The data dictionary containing the object_permission to update
|
||||
existing_object_permission_id: The current object_permission_id from the entity (can be None)
|
||||
prisma_client: The database client
|
||||
|
||||
Returns:
|
||||
Optional[str]: The object_permission_id after the update/creation, or None if no object_permission to process
|
||||
|
||||
Raises:
|
||||
ValueError: If prisma_client is None
|
||||
"""
|
||||
if prisma_client is None:
|
||||
raise ValueError("Prisma client not found")
|
||||
|
||||
#########################################################
|
||||
# Ensure `object_permission` is not added to the data_json
|
||||
# We need to update the entity at the object_permission_id level in the LiteLLM_ObjectPermissionTable
|
||||
#########################################################
|
||||
new_object_permission: Union[dict, str] = data_json.pop("object_permission", None)
|
||||
if new_object_permission is None:
|
||||
return None
|
||||
|
||||
# Lookup existing object permission ID and update that entry
|
||||
object_permission_id_to_use: str = existing_object_permission_id or str(
|
||||
uuid.uuid4()
|
||||
)
|
||||
existing_object_permissions_dict: Dict = {}
|
||||
|
||||
existing_object_permission = (
|
||||
await prisma_client.db.litellm_objectpermissiontable.find_unique(
|
||||
where={"object_permission_id": object_permission_id_to_use},
|
||||
)
|
||||
)
|
||||
|
||||
# Update the object permission
|
||||
if existing_object_permission is not None:
|
||||
existing_object_permissions_dict = existing_object_permission.model_dump(
|
||||
exclude_unset=True, exclude_none=True
|
||||
)
|
||||
|
||||
# Handle string JSON object permission
|
||||
if isinstance(new_object_permission, str):
|
||||
new_object_permission = json.loads(new_object_permission)
|
||||
|
||||
if isinstance(new_object_permission, dict):
|
||||
existing_object_permissions_dict.update(new_object_permission)
|
||||
|
||||
#########################################################
|
||||
# Commit the update to the LiteLLM_ObjectPermissionTable
|
||||
#########################################################
|
||||
created_object_permission_row = (
|
||||
await prisma_client.db.litellm_objectpermissiontable.upsert(
|
||||
where={"object_permission_id": object_permission_id_to_use},
|
||||
data={
|
||||
"create": existing_object_permissions_dict,
|
||||
"update": existing_object_permissions_dict,
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
f"created_object_permission_row: {created_object_permission_row}"
|
||||
)
|
||||
|
||||
return created_object_permission_row.object_permission_id
|
||||
@@ -169,8 +169,8 @@ class VectorStoreRegistry:
|
||||
Only add the vector store if it is not already in the registry
|
||||
"""
|
||||
vector_store_id = vector_store.get("vector_store_id")
|
||||
for vector_store in self.vector_stores:
|
||||
if vector_store.get("vector_store_id") == vector_store_id:
|
||||
for _vector_store in self.vector_stores:
|
||||
if _vector_store.get("vector_store_id") == vector_store_id:
|
||||
return
|
||||
self.vector_stores.append(vector_store)
|
||||
|
||||
@@ -209,3 +209,18 @@ class VectorStoreRegistry:
|
||||
)
|
||||
vector_stores_from_db.append(_litellm_managed_vector_store)
|
||||
return vector_stores_from_db
|
||||
|
||||
def get_credentials_for_vector_store(self, vector_store_id: str) -> Dict[str, Any]:
|
||||
"""
|
||||
Get the credentials for a vector store
|
||||
|
||||
Returns a dictionary of unpacked credentials for the vector store to use for the request
|
||||
"""
|
||||
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
|
||||
|
||||
for vector_store in self.vector_stores:
|
||||
if vector_store.get("vector_store_id") == vector_store_id:
|
||||
credentials = vector_store.get("litellm_credential_name")
|
||||
if credentials:
|
||||
return CredentialAccessor.get_credential_values(credentials)
|
||||
return {}
|
||||
|
||||
@@ -388,12 +388,6 @@ async def test_key_update_object_permissions_no_existing_permission(monkeypatch)
|
||||
# Verify new object_permission_id was set
|
||||
assert "object_permission" not in result
|
||||
assert result["object_permission_id"] == "new_perm_id_456"
|
||||
|
||||
# Verify find_unique was called with None
|
||||
mock_prisma_client.db.litellm_objectpermissiontable.find_unique.assert_called_once_with(
|
||||
where={"object_permission_id": None}
|
||||
)
|
||||
|
||||
# Verify upsert was called to create new record
|
||||
mock_prisma_client.db.litellm_objectpermissiontable.upsert.assert_called_once()
|
||||
|
||||
|
||||
@@ -159,11 +159,6 @@ async def test_organization_update_object_permissions_no_existing_permission(
|
||||
assert "object_permission" not in result
|
||||
assert result["object_permission_id"] == "new_perm_id_456"
|
||||
|
||||
# Verify find_unique was called with None
|
||||
mock_prisma_client.db.litellm_objectpermissiontable.find_unique.assert_called_once_with(
|
||||
where={"object_permission_id": None}
|
||||
)
|
||||
|
||||
# Verify upsert was called to create new record
|
||||
mock_prisma_client.db.litellm_objectpermissiontable.upsert.assert_called_once()
|
||||
|
||||
|
||||
@@ -443,11 +443,6 @@ async def test_team_update_object_permissions_no_existing_permission(monkeypatch
|
||||
assert "object_permission" not in result
|
||||
assert result["object_permission_id"] == "new_perm_id_456"
|
||||
|
||||
# Verify find_unique was called with None
|
||||
mock_prisma_client.db.litellm_objectpermissiontable.find_unique.assert_called_once_with(
|
||||
where={"object_permission_id": None}
|
||||
)
|
||||
|
||||
# Verify upsert was called to create new record
|
||||
mock_prisma_client.db.litellm_objectpermissiontable.upsert.assert_called_once()
|
||||
|
||||
|
||||
@@ -0,0 +1,115 @@
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import respx
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../..")
|
||||
) # Adds the parent directory to the system path
|
||||
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from litellm.types.vector_stores import LiteLLM_ManagedVectorStore
|
||||
from litellm.vector_stores.vector_store_registry import VectorStoreRegistry
|
||||
|
||||
|
||||
def test_get_credentials_for_vector_store():
|
||||
"""Test that get_credentials_for_vector_store returns correct credentials"""
|
||||
# Create test vector stores
|
||||
vector_store_1 = LiteLLM_ManagedVectorStore(
|
||||
vector_store_id="test_id_1",
|
||||
custom_llm_provider="openai",
|
||||
vector_store_name="test_store_1",
|
||||
litellm_credential_name="test_creds_1",
|
||||
created_at=datetime.now(timezone.utc),
|
||||
updated_at=datetime.now(timezone.utc),
|
||||
)
|
||||
|
||||
vector_store_2 = LiteLLM_ManagedVectorStore(
|
||||
vector_store_id="test_id_2",
|
||||
custom_llm_provider="bedrockc",
|
||||
vector_store_name="test_store_2",
|
||||
litellm_credential_name="test_creds_2",
|
||||
created_at=datetime.now(timezone.utc),
|
||||
updated_at=datetime.now(timezone.utc),
|
||||
)
|
||||
|
||||
# Create registry with vector stores
|
||||
registry = VectorStoreRegistry([vector_store_1, vector_store_2])
|
||||
|
||||
# Mock CredentialAccessor.get_credential_values
|
||||
with patch(
|
||||
"litellm.litellm_core_utils.credential_accessor.CredentialAccessor.get_credential_values"
|
||||
) as mock_get_creds:
|
||||
mock_get_creds.return_value = {"api_key": "test_key_1", "env": "test"}
|
||||
|
||||
# Test getting credentials for existing vector store
|
||||
result = registry.get_credentials_for_vector_store("test_id_1")
|
||||
|
||||
assert result == {"api_key": "test_key_1", "env": "test"}
|
||||
mock_get_creds.assert_called_once_with("test_creds_1")
|
||||
|
||||
# Test getting credentials for non-existent vector store
|
||||
result = registry.get_credentials_for_vector_store("non_existent_id")
|
||||
assert result == {}
|
||||
|
||||
|
||||
def test_add_vector_store_to_registry():
|
||||
"""Test that add_vector_store_to_registry adds vector store correctly when there are pre-existing stores"""
|
||||
# Create pre-existing vector stores
|
||||
existing_store_1 = LiteLLM_ManagedVectorStore(
|
||||
vector_store_id="existing_id_1",
|
||||
custom_llm_provider="openai",
|
||||
vector_store_name="existing_store_1",
|
||||
created_at=datetime.now(timezone.utc),
|
||||
updated_at=datetime.now(timezone.utc),
|
||||
)
|
||||
|
||||
existing_store_2 = LiteLLM_ManagedVectorStore(
|
||||
vector_store_id="existing_id_2",
|
||||
custom_llm_provider="openai",
|
||||
vector_store_name="existing_store_2",
|
||||
created_at=datetime.now(timezone.utc),
|
||||
updated_at=datetime.now(timezone.utc),
|
||||
)
|
||||
|
||||
# Create registry with pre-existing stores
|
||||
registry = VectorStoreRegistry([existing_store_1, existing_store_2])
|
||||
assert len(registry.vector_stores) == 2
|
||||
|
||||
# Add a new vector store
|
||||
new_store = LiteLLM_ManagedVectorStore(
|
||||
vector_store_id="new_id",
|
||||
custom_llm_provider="bedrock",
|
||||
vector_store_name="new_store",
|
||||
created_at=datetime.now(timezone.utc),
|
||||
updated_at=datetime.now(timezone.utc),
|
||||
)
|
||||
|
||||
registry.add_vector_store_to_registry(new_store)
|
||||
|
||||
# Verify new store was added
|
||||
assert len(registry.vector_stores) == 3
|
||||
assert registry.vector_stores[2]["vector_store_id"] == "new_id"
|
||||
assert registry.vector_stores[2]["vector_store_name"] == "new_store"
|
||||
|
||||
# Try to add duplicate - should not be added
|
||||
duplicate_store = LiteLLM_ManagedVectorStore(
|
||||
vector_store_id="existing_id_1", # Same ID as existing store
|
||||
custom_llm_provider="different_provider",
|
||||
vector_store_name="duplicate_store",
|
||||
created_at=datetime.now(timezone.utc),
|
||||
updated_at=datetime.now(timezone.utc),
|
||||
)
|
||||
|
||||
registry.add_vector_store_to_registry(duplicate_store)
|
||||
|
||||
# Verify duplicate was not added
|
||||
assert len(registry.vector_stores) == 3
|
||||
# Original store should still be there unchanged
|
||||
assert registry.vector_stores[0]["vector_store_name"] == "existing_store_1"
|
||||
Reference in New Issue
Block a user