[Feat] New Vector Store - PG Vector (#12667)

* add PGVectorStoreConfig

* add PGVectorStoreConfig

* test_environment_variable_support

* fix code QA check

* rename test

* add PG vector img

* allow adding vector stores

* add pg vector

* add vector store

* TestPGVectorStoreConfig

* TestPGVectorStoreConfig
This commit is contained in:
Ishaan Jaff 2025-07-16 18:17:05 -07:00 • committed by GitHub
parent b2080ec9af
commit e5f0a8477b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 418 additions and 35 deletions

View file

@ -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"

View file

@ -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

View file

@ -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

View file

@ -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"

View file

@ -0,0 +1 @@
<svg xmlns="http://www.w3.org/2000/svg" height="64" viewBox="0 0 25.6 25.6" width="64"><style><![CDATA[.B{stroke-linecap:round}.C{stroke-linejoin:round}.D{stroke-linejoin:miter}.E{stroke-width:.716}]]></style><g fill="none" stroke="#fff"><path d="M18.983 18.636c.163-1.357.114-1.555 1.124-1.336l.257.023c.777.035 1.793-.125 2.4-.402 1.285-.596 2.047-1.592.78-1.33-2.89.596-3.1-.383-3.1-.383 3.053-4.53 4.33-10.28 3.227-11.687-3.004-3.84-8.205-2.024-8.292-1.976l-.028.005c-.57-.12-1.2-.19-1.93-.2-1.308-.02-2.3.343-3.054.914 0 0-9.277-3.822-8.846 4.807.092 1.836 2.63 13.9 5.66 10.25C8.29 15.987 9.36 14.86 9.36 14.86c.53.353 1.167.533 1.834.468l.052-.044a2.01 2.01 0 0 0 .021.518c-.78.872-.55 1.025-2.11 1.346-1.578.325-.65.904-.046 1.056.734.184 2.432.444 3.58-1.162l-.046.183c.306.245.285 1.76.33 2.842s.116 2.093.337 2.688.48 2.13 2.53 1.7c1.713-.367 3.023-.896 3.143-5.81" fill="#000" stroke="#000" stroke-linecap="butt" stroke-width="2.149" class="D"/><path d="M23.535 15.6c-2.89.596-3.1-.383-3.1-.383 3.053-4.53 4.33-10.28 3.228-11.687-3.004-3.84-8.205-2.023-8.292-1.976l-.028.005a10.31 10.31 0 0 0-1.929-.201c-1.308-.02-2.3.343-3.054.914 0 0-9.278-3.822-8.846 4.807.092 1.836 2.63 13.9 5.66 10.25C8.29 15.987 9.36 14.86 9.36 14.86c.53.353 1.167.533 1.834.468l.052-.044a2.02 2.02 0 0 0 .021.518c-.78.872-.55 1.025-2.11 1.346-1.578.325-.65.904-.046 1.056.734.184 2.432.444 3.58-1.162l-.046.183c.306.245.52 1.593.484 2.815s-.06 2.06.18 2.716.48 2.13 2.53 1.7c1.713-.367 2.6-1.32 2.725-2.906.088-1.128.286-.962.3-1.97l.16-.478c.183-1.53.03-2.023 1.085-1.793l.257.023c.777.035 1.794-.125 2.39-.402 1.285-.596 2.047-1.592.78-1.33z" fill="#336791" stroke="none"/><g class="E"><g class="B"><path d="M12.814 16.467c-.08 2.846.02 5.712.298 6.4s.875 2.05 2.926 1.612c1.713-.367 2.337-1.078 2.607-2.647l.633-5.017M10.356 2.2S1.072-1.596 1.504 7.033c.092 1.836 2.63 13.9 5.66 10.25C8.27 15.95 9.27 14.907 9.27 14.907m6.1-13.4c-.32.1 5.164-2.005 8.282 1.978 1.1 1.407-.175 7.157-3.228 11.687" class="C"/><path d="M20.425 15.17s.2.98 3.1.382c1.267-.262.504.734-.78 1.33-1.054.49-3.418.615-3.457-.06-.1-1.745 1.244-1.215 1.147-1.652-.088-.394-.69-.78-1.086-1.744-.347-.84-4.76-7.29 1.224-6.333.22-.045-1.56-5.7-7.16-5.782S7.99 8.196 7.99 8.196" stroke-linejoin="bevel"/></g><g class="C"><path d="M11.247 15.768c-.78.872-.55 1.025-2.11 1.346-1.578.325-.65.904-.046 1.056.734.184 2.432.444 3.58-1.163.35-.49-.002-1.27-.482-1.468-.232-.096-.542-.216-.94.23z"/><path d="M11.196 15.753c-.08-.513.168-1.122.433-1.836.398-1.07 1.316-2.14.582-5.537-.547-2.53-4.22-.527-4.22-.184s.166 1.74-.06 3.365c-.297 2.122 1.35 3.916 3.246 3.733" class="B"/></g></g><g fill="#fff" class="D"><path d="M10.322 8.145c-.017.117.215.43.516.472s.558-.202.575-.32-.215-.246-.516-.288-.56.02-.575.136z" stroke-width=".239"/><path d="M19.486 7.906c.016.117-.215.43-.516.472s-.56-.202-.575-.32.215-.246.516-.288.56.02.575.136z" stroke-width=".119"/></g><path d="M20.562 7.095c.05.92-.198 1.545-.23 2.524-.046 1.422.678 3.05-.413 4.68" class="B C E"/></g></svg>

After

Width:  |  Height:  |  Size: 3 KiB

View file

@ -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<VectorStoreFormProps> = ({
}) => {
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<VectorStoreFormProps> = ({
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<VectorStoreFormProps> = ({
const handleCancel = () => {
form.resetFields();
setMetadataJson("{}");
setSelectedProvider("bedrock");
onCancel();
};
return (
<Modal
title="Create New Vector Store"
title="Add New Vector Store"
visible={isVisible}
width={800}
width={1000}
footer={null}
onCancel={handleCancel}
>
@ -99,38 +113,56 @@ const VectorStoreForm: React.FC<VectorStoreFormProps> = ({
rules={[{ required: true, message: "Please select a provider" }]}
initialValue="bedrock"
>
<Select>
{Object.entries(Providers).map(([providerEnum, providerDisplayName]) => {
// Currently only showing Bedrock since it's the only supported provider
if (providerEnum === 'Bedrock') {
return (
<Select.Option key={providerEnum} value={provider_map[providerEnum]}>
<div className="flex items-center space-x-2">
<img
src={providerLogoMap[providerDisplayName]}
alt={`${providerEnum} logo`}
className="w-5 h-5"
onError={(e) => {
// 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);
}
}}
/>
<span>{providerDisplayName}</span>
</div>
</Select.Option>
);
}
return null;
<Select onChange={(value) => setSelectedProvider(value)}>
{Object.entries(VectorStoreProviders).map(([providerEnum, providerDisplayName]) => {
return (
<Select.Option key={providerEnum} value={vectorStoreProviderMap[providerEnum]}>
<div className="flex items-center space-x-2">
<img
src={vectorStoreProviderLogoMap[providerDisplayName]}
alt={`${providerEnum} logo`}
className="w-5 h-5"
onError={(e) => {
// 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);
}
}}
/>
<span>{providerDisplayName}</span>
</div>
</Select.Option>
);
})}
</Select>
</Form.Item>
{/* PG Vector Setup Instructions */}
{selectedProvider === "pg_vector" && (
<Alert
message="PG Vector Setup Required"
description={
<div>
<p>LiteLLM provides a server to connect to PG Vector. To use this provider:</p>
<ol style={{ marginLeft: '16px', marginTop: '8px' }}>
<li>Deploy the litellm-pgvector server from: <a href="https://github.com/BerriAI/litellm-pgvector" target="_blank" rel="noopener noreferrer">https://github.com/BerriAI/litellm-pgvector</a></li>
<li>Configure your PostgreSQL database with pgvector extension</li>
<li>Start the server and note the API base URL and API key</li>
<li>Enter those details in the fields below</li>
</ol>
</div>
}
type="info"
showIcon
style={{ marginBottom: '16px' }}
/>
)}
<Form.Item
label={
<span>
@ -146,6 +178,28 @@ const VectorStoreForm: React.FC<VectorStoreFormProps> = ({
<TextInput />
</Form.Item>
{/* Provider-specific fields */}
{getProviderSpecificFields(selectedProvider).map((field: VectorStoreFieldConfig) => (
<Form.Item
key={field.name}
label={
<span>
{field.label}{' '}
<Tooltip title={field.tooltip}>
<InfoCircleOutlined style={{ marginLeft: '4px' }} />
</Tooltip>
</span>
}
name={field.name}
rules={field.required ? [{ required: true, message: `Please input the ${field.label.toLowerCase()}` }] : []}
>
<TextInput
type={field.type || "text"}
placeholder={field.placeholder}
/>
</Form.Item>
))}
<Form.Item
label={
<span>

View file

@ -121,7 +121,7 @@ const VectorStoreManagement: React.FC<VectorStoreProps> = ({
className="mb-4"
onClick={() => setIsCreateModalVisible(true)}
>
+ Create Vector Store
+ Add Vector Store
</TremorButton>
<Grid numItems={1} className="gap-2 pt-2 pb-2 h-[75vh] w-full mt-2">

View file

@ -0,0 +1,74 @@
export enum VectorStoreProviders {
Bedrock = "Amazon Bedrock",
PgVector = "PostgreSQL pgvector (LiteLLM Connector)"
}
export const vectorStoreProviderMap: Record<string, string> = {
Bedrock: "bedrock",
PgVector: "pg_vector"
};
const asset_logos_folder = '/assets/logos/';
export const vectorStoreProviderLogoMap: Record<string, string> = {
[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<string, VectorStoreFieldConfig[]> = {
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] || [];
};