diff --git a/litellm/llms/pg_vector/vector_stores/transformation.py b/litellm/llms/pg_vector/vector_stores/transformation.py new file mode 100644 index 00000000000..1bc03a6f2f8 --- /dev/null +++ b/litellm/llms/pg_vector/vector_stores/transformation.py @@ -0,0 +1,69 @@ +from typing import Optional + +from litellm.llms.openai.vector_stores.transformation import OpenAIVectorStoreConfig +from litellm.secret_managers.main import get_secret_str +from litellm.types.router import GenericLiteLLMParams + + +class PGVectorStoreConfig(OpenAIVectorStoreConfig): + """ + PG Vector Store configuration that inherits from OpenAI since it's OpenAI-compatible. + + LiteLLM Provides an OpenAI Compatible Server to connect to PG Vector. + + https://github.com/BerriAI/litellm-pgvector + + You just need to connect litellm proxy to this deployed server. + + Requires: + - api_base: The base URL for the PG vector service + - api_key: API key for authentication with the PG vector service + """ + + def validate_environment( + self, headers: dict, litellm_params: Optional[GenericLiteLLMParams] + ) -> dict: + """ + Validate environment and set headers for PG vector service authentication + """ + litellm_params = litellm_params or GenericLiteLLMParams() + + # Get API key from various sources + api_key = ( + litellm_params.api_key + or get_secret_str("PG_VECTOR_API_KEY") + ) + + if not api_key: + raise ValueError("PG Vector API key is required. Set PG_VECTOR_API_KEY environment variable or pass api_key in litellm_params.") + + headers.update( + { + "Authorization": f"Bearer {api_key}", + "Content-Type": "application/json", + } + ) + + return headers + + def get_complete_url( + self, + api_base: Optional[str], + litellm_params: dict, + ) -> str: + """ + Get the complete URL for PG vector service endpoints + """ + # Get API base from various sources + api_base = ( + api_base + or get_secret_str("PG_VECTOR_API_BASE") + ) + + if not api_base: + raise ValueError("PG Vector API base URL is required. Set PG_VECTOR_API_BASE environment variable or pass api_base in litellm_params.") + + # Remove trailing slashes + api_base = api_base.rstrip("/") + + return f"{api_base}/vector_stores" \ No newline at end of file diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 74608a643a4..9ad3005dfdb 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -2309,12 +2309,12 @@ class LlmProviders(str, Enum): SNOWFLAKE = "snowflake" LLAMA = "meta_llama" NSCALE = "nscale" + PG_VECTOR = "pg_vector" # Create a set of all provider values for quick lookup LlmProvidersSet = {provider.value for provider in LlmProviders} - class LiteLLMLoggingBaseClass: """ Base class for logging pre and post call diff --git a/litellm/utils.py b/litellm/utils.py index fb87a4551cc..f92d14b5dd7 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -7101,6 +7101,12 @@ class ProviderConfigManager: ) return VertexVectorStoreConfig() + elif litellm.LlmProviders.PG_VECTOR == provider: + from litellm.llms.pg_vector.vector_stores.transformation import ( + PGVectorStoreConfig, + ) + + return PGVectorStoreConfig() return None @staticmethod diff --git a/tests/test_litellm/llms/pg_vector/vector_stores/test_pg_vector_transformation.py b/tests/test_litellm/llms/pg_vector/vector_stores/test_pg_vector_transformation.py new file mode 100644 index 00000000000..644980928ae --- /dev/null +++ b/tests/test_litellm/llms/pg_vector/vector_stores/test_pg_vector_transformation.py @@ -0,0 +1,179 @@ +""" +Unit tests for PG Vector Store transformation. + +This test file mirrors litellm/llms/pg_vector/vector_stores/transformation.py +and contains mocked tests for the PGVectorStoreConfig class. +""" + +from unittest.mock import Mock + +import pytest + +from litellm.llms.pg_vector.vector_stores.transformation import PGVectorStoreConfig +from litellm.types.router import GenericLiteLLMParams + + +class TestPGVectorStoreConfig: + """Test the PG Vector Store transformation configuration.""" + + def test_validate_environment_with_api_key_in_params(self): + """ + Test that validate_environment works when api_key is provided in litellm_params. + + This test validates that API key from params is correctly set in headers. + """ + config = PGVectorStoreConfig() + litellm_params = GenericLiteLLMParams(api_key="test_pg_vector_key_123") + headers = {} + + result_headers = config.validate_environment(headers, litellm_params) + + assert "Authorization" in result_headers + assert result_headers["Authorization"] == "Bearer test_pg_vector_key_123" + assert result_headers["Content-Type"] == "application/json" + + def test_validate_environment_missing_api_key(self): + """ + Test that validate_environment raises ValueError when no API key is provided. + + This test validates that proper error handling occurs for missing credentials. + """ + config = PGVectorStoreConfig() + litellm_params = GenericLiteLLMParams() + headers = {} + + with pytest.raises(ValueError) as exc_info: + config.validate_environment(headers, litellm_params) + + assert "PG Vector API key is required" in str(exc_info.value) + + def test_get_complete_url_with_api_base(self): + """ + Test that get_complete_url correctly formats the URL with api_base. + + This test validates URL construction for PG Vector endpoints. + """ + config = PGVectorStoreConfig() + api_base = "https://my-pg-vector-service.example.com" + litellm_params = {} + + result_url = config.get_complete_url(api_base, litellm_params) + + assert result_url == "https://my-pg-vector-service.example.com/vector_stores" + + def test_get_complete_url_removes_trailing_slashes(self): + """ + Test that get_complete_url handles trailing slashes correctly. + + This test validates that URLs are normalized properly. + """ + config = PGVectorStoreConfig() + api_base = "https://my-pg-vector-service.example.com/" + litellm_params = {} + + result_url = config.get_complete_url(api_base, litellm_params) + + assert result_url == "https://my-pg-vector-service.example.com/vector_stores" + + def test_get_complete_url_missing_api_base(self): + """ + Test that get_complete_url raises ValueError when no API base is provided. + + This test validates that proper error handling occurs for missing API base. + """ + config = PGVectorStoreConfig() + litellm_params = {} + + with pytest.raises(ValueError) as exc_info: + config.get_complete_url(None, litellm_params) + + assert "PG Vector API base URL is required" in str(exc_info.value) + + def test_inheritance_from_openai_config(self): + """ + Test that PGVectorStoreConfig correctly inherits from OpenAIVectorStoreConfig. + + This test validates that PG Vector config inherits OpenAI-compatible methods. + """ + from litellm.llms.openai.vector_stores.transformation import ( + OpenAIVectorStoreConfig, + ) + + config = PGVectorStoreConfig() + + # Test that it's an instance of the parent class + assert isinstance(config, OpenAIVectorStoreConfig) + + # Test that inherited methods are available + assert hasattr(config, 'transform_search_vector_store_request') + assert hasattr(config, 'transform_search_vector_store_response') + assert hasattr(config, 'transform_create_vector_store_request') + assert hasattr(config, 'transform_create_vector_store_response') + + def test_openai_compatible_methods_available(self): + """ + Test that OpenAI-compatible transformation methods are available. + + Since PG Vector is OpenAI-compatible, it should inherit all transformation methods. + """ + config = PGVectorStoreConfig() + + # Test that transformation methods are callable + assert callable(getattr(config, 'transform_search_vector_store_request', None)) + assert callable(getattr(config, 'transform_search_vector_store_response', None)) + assert callable(getattr(config, 'transform_create_vector_store_request', None)) + assert callable(getattr(config, 'transform_create_vector_store_response', None)) + + def test_config_methods_with_mock_data(self): + """ + Test configuration with mock data to ensure basic functionality. + + This test validates that the config works with typical parameters. + """ + config = PGVectorStoreConfig() + + # Test with valid parameters + litellm_params = GenericLiteLLMParams(api_key="test_key") + headers = config.validate_environment({}, litellm_params) + url = config.get_complete_url("https://example.com", {}) + + # Verify results + assert headers["Authorization"] == "Bearer test_key" + assert url == "https://example.com/vector_stores" + + @pytest.mark.serial + def test_environment_variable_support(self): + """ + Test that environment variables are supported for configuration. + + This test validates that the config properly reads from environment variables. + """ + import os + from unittest.mock import patch + + config = PGVectorStoreConfig() + + # Test API key from environment variable + with patch.dict(os.environ, {'PG_VECTOR_API_KEY': 'env_api_key_123'}): + litellm_params = GenericLiteLLMParams() # No API key in params + + headers = config.validate_environment({}, litellm_params) + + assert headers["Authorization"] == "Bearer env_api_key_123" + assert headers["Content-Type"] == "application/json" + + # Test API base from environment variable + with patch.dict(os.environ, {'PG_VECTOR_API_BASE': 'https://env-pg-vector.example.com'}): + url = config.get_complete_url(None, {}) + + assert url == "https://env-pg-vector.example.com/vector_stores" + + # Test that params take precedence over environment variables + with patch.dict(os.environ, {'PG_VECTOR_API_KEY': 'env_key'}): + litellm_params = GenericLiteLLMParams(api_key="param_key") + + headers = config.validate_environment({}, litellm_params) + + # Param key should take precedence over environment variable + assert headers["Authorization"] == "Bearer param_key" + assert headers["Content-Type"] == "application/json" \ No newline at end of file diff --git a/ui/litellm-dashboard/public/assets/logos/postgresql.svg b/ui/litellm-dashboard/public/assets/logos/postgresql.svg new file mode 100644 index 00000000000..7fed68bcd33 --- /dev/null +++ b/ui/litellm-dashboard/public/assets/logos/postgresql.svg @@ -0,0 +1 @@ + \ No newline at end of file diff --git a/ui/litellm-dashboard/src/components/vector_store_management/VectorStoreForm.tsx b/ui/litellm-dashboard/src/components/vector_store_management/VectorStoreForm.tsx index 1e9b5024c41..6c6db2a954e 100644 --- a/ui/litellm-dashboard/src/components/vector_store_management/VectorStoreForm.tsx +++ b/ui/litellm-dashboard/src/components/vector_store_management/VectorStoreForm.tsx @@ -12,10 +12,11 @@ import { message, Tooltip, Input, + Alert, } from "antd"; import { InfoCircleOutlined } from '@ant-design/icons'; import { CredentialItem, vectorStoreCreateCall } from "../networking"; -import { Providers, providerLogoMap, provider_map } from "../provider_info_helpers"; +import { VectorStoreProviders, vectorStoreProviderLogoMap, vectorStoreProviderMap, getProviderSpecificFields, VectorStoreFieldConfig } from "../vector_store_providers"; interface VectorStoreFormProps { isVisible: boolean; @@ -34,6 +35,7 @@ const VectorStoreForm: React.FC = ({ }) => { const [form] = Form.useForm(); const [metadataJson, setMetadataJson] = useState("{}"); + const [selectedProvider, setSelectedProvider] = useState("bedrock"); const handleCreate = async (formValues: any) => { if (!accessToken) return; @@ -47,14 +49,25 @@ const VectorStoreForm: React.FC = ({ return; } - await vectorStoreCreateCall(accessToken, { + // Prepare the payload with provider-specific fields + const payload: any = { vector_store_id: formValues.vector_store_id, custom_llm_provider: formValues.custom_llm_provider, vector_store_name: formValues.vector_store_name, vector_store_description: formValues.vector_store_description, vector_store_metadata: metadata, litellm_credential_name: formValues.litellm_credential_name, + }; + + // Add provider-specific fields dynamically + const providerFields = getProviderSpecificFields(formValues.custom_llm_provider); + providerFields.forEach(field => { + if (formValues[field.name]) { + payload[field.name] = formValues[field.name]; + } }); + + await vectorStoreCreateCall(accessToken, payload); message.success("Vector store created successfully"); form.resetFields(); setMetadataJson("{}"); @@ -68,14 +81,15 @@ const VectorStoreForm: React.FC = ({ const handleCancel = () => { form.resetFields(); setMetadataJson("{}"); + setSelectedProvider("bedrock"); onCancel(); }; return ( @@ -99,38 +113,56 @@ const VectorStoreForm: React.FC = ({ rules={[{ required: true, message: "Please select a provider" }]} initialValue="bedrock" > - setSelectedProvider(value)}> + {Object.entries(VectorStoreProviders).map(([providerEnum, providerDisplayName]) => { + return ( + +
+ {`${providerEnum} { + // Create a div with provider initial as fallback + const target = e.target as HTMLImageElement; + const parent = target.parentElement; + if (parent) { + const fallbackDiv = document.createElement('div'); + fallbackDiv.className = 'w-5 h-5 rounded-full bg-gray-200 flex items-center justify-center text-xs'; + fallbackDiv.textContent = providerDisplayName.charAt(0); + parent.replaceChild(fallbackDiv, target); + } + }} + /> + {providerDisplayName} +
+
+ ); })} + + {/* PG Vector Setup Instructions */} + {selectedProvider === "pg_vector" && ( + +

LiteLLM provides a server to connect to PG Vector. To use this provider:

+
    +
  1. Deploy the litellm-pgvector server from: https://github.com/BerriAI/litellm-pgvector
  2. +
  3. Configure your PostgreSQL database with pgvector extension
  4. +
  5. Start the server and note the API base URL and API key
  6. +
  7. Enter those details in the fields below
  8. +
+ + } + type="info" + showIcon + style={{ marginBottom: '16px' }} + /> + )} + @@ -146,6 +178,28 @@ const VectorStoreForm: React.FC = ({ + {/* Provider-specific fields */} + {getProviderSpecificFields(selectedProvider).map((field: VectorStoreFieldConfig) => ( + + {field.label}{' '} + + + + + } + name={field.name} + rules={field.required ? [{ required: true, message: `Please input the ${field.label.toLowerCase()}` }] : []} + > + + + ))} + diff --git a/ui/litellm-dashboard/src/components/vector_store_management/index.tsx b/ui/litellm-dashboard/src/components/vector_store_management/index.tsx index d01fd1417af..defe621eab2 100644 --- a/ui/litellm-dashboard/src/components/vector_store_management/index.tsx +++ b/ui/litellm-dashboard/src/components/vector_store_management/index.tsx @@ -121,7 +121,7 @@ const VectorStoreManagement: React.FC = ({ className="mb-4" onClick={() => setIsCreateModalVisible(true)} > - + Create Vector Store + + Add Vector Store diff --git a/ui/litellm-dashboard/src/components/vector_store_providers.tsx b/ui/litellm-dashboard/src/components/vector_store_providers.tsx new file mode 100644 index 00000000000..9cc97b99478 --- /dev/null +++ b/ui/litellm-dashboard/src/components/vector_store_providers.tsx @@ -0,0 +1,74 @@ +export enum VectorStoreProviders { + Bedrock = "Amazon Bedrock", + PgVector = "PostgreSQL pgvector (LiteLLM Connector)" +} + +export const vectorStoreProviderMap: Record = { + Bedrock: "bedrock", + PgVector: "pg_vector" +}; + +const asset_logos_folder = '/assets/logos/'; + +export const vectorStoreProviderLogoMap: Record = { + [VectorStoreProviders.Bedrock]: `${asset_logos_folder}bedrock.svg`, + [VectorStoreProviders.PgVector]: `${asset_logos_folder}postgresql.svg`, // Fallback to a generic database icon if needed +}; + +// Define field types for provider-specific configurations +export interface VectorStoreFieldConfig { + name: string; + label: string; + tooltip: string; + placeholder?: string; + required: boolean; + type?: 'text' | 'password'; +} + +// Provider-specific field configurations +export const vectorStoreProviderFields: Record = { + bedrock: [], + pg_vector: [ + { + name: "api_base", + label: "API Base", + tooltip: "Enter the base URL of your deployed litellm-pgvector server (e.g., http://your-server:8000)", + placeholder: "http://your-deployed-server:8000", + required: true, + type: "text" + }, + { + name: "api_key", + label: "API Key", + tooltip: "Enter the API key from your deployed litellm-pgvector server", + placeholder: "your-deployed-api-key", + required: true, + type: "password" + } + ] +}; + +export const getVectorStoreProviderLogoAndName = (providerValue: string): { logo: string, displayName: string } => { + if (!providerValue) { + return { logo: "", displayName: "-" }; + } + + // Find the enum key by matching vectorStoreProviderMap values + const enumKey = Object.keys(vectorStoreProviderMap).find( + key => vectorStoreProviderMap[key].toLowerCase() === providerValue.toLowerCase() + ); + + if (!enumKey) { + return { logo: "", displayName: providerValue }; + } + + // Get the display name from VectorStoreProviders enum and logo from map + const displayName = VectorStoreProviders[enumKey as keyof typeof VectorStoreProviders]; + const logo = vectorStoreProviderLogoMap[displayName as keyof typeof vectorStoreProviderLogoMap]; + + return { logo, displayName }; +}; + +export const getProviderSpecificFields = (providerValue: string): VectorStoreFieldConfig[] => { + return vectorStoreProviderFields[providerValue] || []; +}; \ No newline at end of file