mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
feat: _save_vector_store_to_db_from_rag_ingest
This commit is contained in:
parent
cec1a3c858
commit
75eec14691
2 changed files with 216 additions and 58 deletions
|
|
@ -26,6 +26,75 @@ from litellm.proxy.common_utils.http_parsing_utils import (
|
|||
router = APIRouter()
|
||||
|
||||
|
||||
async def _save_vector_store_to_db_from_rag_ingest(
|
||||
response: Any,
|
||||
ingest_options: Dict[str, Any],
|
||||
prisma_client,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> None:
|
||||
"""
|
||||
Helper function to save a newly created vector store from RAG ingest to the database.
|
||||
|
||||
This function:
|
||||
- Extracts vector store ID and config from the ingest response
|
||||
- Checks if the vector store already exists in the database
|
||||
- Creates a new database entry if it doesn't exist
|
||||
- Adds the vector store to the registry
|
||||
|
||||
Args:
|
||||
response: The response from litellm.aingest()
|
||||
ingest_options: The ingest options containing vector store config
|
||||
prisma_client: The Prisma database client
|
||||
user_api_key_dict: User API key authentication info
|
||||
"""
|
||||
from litellm.proxy.vector_store_endpoints.management_endpoints import (
|
||||
create_vector_store_in_db,
|
||||
)
|
||||
|
||||
vector_store_id = response.get("vector_store_id")
|
||||
if vector_store_id is None or not isinstance(vector_store_id, str):
|
||||
verbose_proxy_logger.warning(
|
||||
"Vector store ID is None or not a string, skipping database save"
|
||||
)
|
||||
return
|
||||
|
||||
vector_store_config = ingest_options.get("vector_store", {})
|
||||
custom_llm_provider = vector_store_config.get("custom_llm_provider")
|
||||
|
||||
try:
|
||||
# Check if vector store already exists in database
|
||||
existing_vector_store = (
|
||||
await prisma_client.db.litellm_managedvectorstorestable.find_unique(
|
||||
where={"vector_store_id": vector_store_id}
|
||||
)
|
||||
)
|
||||
|
||||
# Only create if it doesn't exist
|
||||
if existing_vector_store is None:
|
||||
verbose_proxy_logger.info(
|
||||
f"Saving newly created vector store {vector_store_id} to database"
|
||||
)
|
||||
|
||||
await create_vector_store_in_db(
|
||||
vector_store_id=vector_store_id,
|
||||
custom_llm_provider=custom_llm_provider or "openai",
|
||||
prisma_client=prisma_client,
|
||||
vector_store_name=f"RAG Vector Store - {vector_store_id[:8]}",
|
||||
vector_store_description="Created via RAG ingest endpoint",
|
||||
created_by=user_api_key_dict.user_id,
|
||||
updated_by=user_api_key_dict.user_id,
|
||||
)
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
f"Vector store {vector_store_id} saved to database successfully"
|
||||
)
|
||||
except Exception as db_error:
|
||||
# Log the error but don't fail the request since ingestion succeeded
|
||||
verbose_proxy_logger.warning(
|
||||
f"Failed to save vector store {vector_store_id} to database: {db_error}"
|
||||
)
|
||||
|
||||
|
||||
async def parse_rag_ingest_request(
|
||||
request: Request,
|
||||
) -> Tuple[Dict[str, Any], Optional[Tuple[str, bytes, str]], Optional[str], Optional[str]]:
|
||||
|
|
@ -158,9 +227,11 @@ async def rag_ingest(
|
|||
add_litellm_data_to_request,
|
||||
general_settings,
|
||||
llm_router,
|
||||
prisma_client,
|
||||
proxy_config,
|
||||
version,
|
||||
)
|
||||
from litellm.types.vector_stores import LiteLLM_ManagedVectorStore
|
||||
|
||||
try:
|
||||
# Parse request
|
||||
|
|
@ -189,6 +260,20 @@ async def rag_ingest(
|
|||
**request_data,
|
||||
)
|
||||
|
||||
# Save vector store to database if it was newly created and prisma_client is available
|
||||
if (
|
||||
prisma_client is not None
|
||||
and response is not None
|
||||
and isinstance(response, dict)
|
||||
and response.get("vector_store_id")
|
||||
):
|
||||
await _save_vector_store_to_db_from_rag_ingest(
|
||||
response=response,
|
||||
ingest_options=ingest_options,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
return response
|
||||
|
||||
except HTTPException:
|
||||
|
|
|
|||
|
|
@ -133,6 +133,115 @@ async def _resolve_embedding_config_from_db(
|
|||
return None
|
||||
|
||||
|
||||
########################################################
|
||||
# Helper Functions
|
||||
########################################################
|
||||
async def create_vector_store_in_db(
|
||||
vector_store_id: str,
|
||||
custom_llm_provider: str,
|
||||
prisma_client,
|
||||
vector_store_name: Optional[str] = None,
|
||||
vector_store_description: Optional[str] = None,
|
||||
vector_store_metadata: Optional[Dict] = None,
|
||||
litellm_params: Optional[Dict] = None,
|
||||
created_by: Optional[str] = None,
|
||||
updated_by: Optional[str] = None,
|
||||
) -> LiteLLM_ManagedVectorStore:
|
||||
"""
|
||||
Helper function to create a vector store in the database.
|
||||
|
||||
This function handles:
|
||||
- Checking if vector store already exists
|
||||
- Creating the vector store in the database
|
||||
- Adding it to the vector store registry
|
||||
|
||||
Returns:
|
||||
LiteLLM_ManagedVectorStore: The created vector store object
|
||||
|
||||
Raises:
|
||||
HTTPException: If vector store already exists or database error occurs
|
||||
"""
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail="Database not connected")
|
||||
|
||||
# Check if vector store already exists
|
||||
existing_vector_store = (
|
||||
await prisma_client.db.litellm_managedvectorstorestable.find_unique(
|
||||
where={"vector_store_id": vector_store_id}
|
||||
)
|
||||
)
|
||||
if existing_vector_store is not None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Vector store with ID {vector_store_id} already exists",
|
||||
)
|
||||
|
||||
# Prepare data for database
|
||||
data_to_create: Dict[str, Any] = {
|
||||
"vector_store_id": vector_store_id,
|
||||
"custom_llm_provider": custom_llm_provider,
|
||||
}
|
||||
|
||||
if vector_store_name is not None:
|
||||
data_to_create["vector_store_name"] = vector_store_name
|
||||
if vector_store_description is not None:
|
||||
data_to_create["vector_store_description"] = vector_store_description
|
||||
if vector_store_metadata is not None:
|
||||
data_to_create["vector_store_metadata"] = safe_dumps(vector_store_metadata)
|
||||
if created_by is not None:
|
||||
data_to_create["created_by"] = created_by
|
||||
if updated_by is not None:
|
||||
data_to_create["updated_by"] = updated_by
|
||||
|
||||
# Handle litellm_params
|
||||
litellm_params_json: Optional[str] = None
|
||||
if litellm_params:
|
||||
# Auto-resolve embedding config if embedding model is provided but config is not
|
||||
embedding_model = litellm_params.get("litellm_embedding_model")
|
||||
if embedding_model and not litellm_params.get("litellm_embedding_config"):
|
||||
resolved_config = await _resolve_embedding_config_from_db(
|
||||
embedding_model=embedding_model,
|
||||
prisma_client=prisma_client
|
||||
)
|
||||
if resolved_config:
|
||||
litellm_params["litellm_embedding_config"] = resolved_config
|
||||
verbose_proxy_logger.info(
|
||||
f"Auto-resolved embedding config for model {embedding_model}"
|
||||
)
|
||||
|
||||
litellm_params_dict = GenericLiteLLMParams(
|
||||
**litellm_params
|
||||
).model_dump(exclude_none=True)
|
||||
litellm_params_json = safe_dumps(litellm_params_dict)
|
||||
|
||||
data_to_create["litellm_params"] = litellm_params_json
|
||||
|
||||
# Create in database
|
||||
_new_vector_store = (
|
||||
await prisma_client.db.litellm_managedvectorstorestable.create(
|
||||
data=data_to_create
|
||||
)
|
||||
)
|
||||
|
||||
new_vector_store: LiteLLM_ManagedVectorStore = LiteLLM_ManagedVectorStore(
|
||||
**_new_vector_store.model_dump()
|
||||
)
|
||||
|
||||
# Add vector store to registry
|
||||
if litellm.vector_store_registry is not None:
|
||||
litellm.vector_store_registry.add_vector_store_to_registry(
|
||||
vector_store=new_vector_store
|
||||
)
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
f"Vector store {vector_store_id} created in database successfully"
|
||||
)
|
||||
|
||||
return new_vector_store
|
||||
|
||||
|
||||
########################################################
|
||||
# Management Endpoints
|
||||
########################################################
|
||||
|
|
@ -156,71 +265,35 @@ async def new_vector_store(
|
|||
- vector_store_metadata: Optional[Dict] - Additional metadata for the vector store
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail="Database not connected")
|
||||
|
||||
try:
|
||||
# Check if vector store already exists
|
||||
existing_vector_store = (
|
||||
await prisma_client.db.litellm_managedvectorstorestable.find_unique(
|
||||
where={"vector_store_id": vector_store.get("vector_store_id")}
|
||||
)
|
||||
)
|
||||
if existing_vector_store is not None:
|
||||
vector_store_id = vector_store.get("vector_store_id")
|
||||
custom_llm_provider = vector_store.get("custom_llm_provider")
|
||||
|
||||
if not vector_store_id or not custom_llm_provider:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Vector store with ID {vector_store.get('vector_store_id')} already exists",
|
||||
)
|
||||
|
||||
if vector_store.get("vector_store_metadata") is not None:
|
||||
vector_store["vector_store_metadata"] = safe_dumps(
|
||||
vector_store.get("vector_store_metadata")
|
||||
)
|
||||
|
||||
# Safely handle JSON serialization of litellm_params
|
||||
litellm_params_json: Optional[str] = None
|
||||
_input_litellm_params: dict = vector_store.get("litellm_params", {}) or {}
|
||||
if _input_litellm_params is not None:
|
||||
# Auto-resolve embedding config if embedding model is provided but config is not
|
||||
embedding_model = _input_litellm_params.get("litellm_embedding_model")
|
||||
if embedding_model and not _input_litellm_params.get("litellm_embedding_config"):
|
||||
resolved_config = await _resolve_embedding_config_from_db(
|
||||
embedding_model=embedding_model,
|
||||
prisma_client=prisma_client
|
||||
)
|
||||
if resolved_config:
|
||||
_input_litellm_params["litellm_embedding_config"] = resolved_config
|
||||
verbose_proxy_logger.info(
|
||||
f"Auto-resolved embedding config for model {embedding_model}"
|
||||
)
|
||||
|
||||
litellm_params_dict = GenericLiteLLMParams(
|
||||
**_input_litellm_params
|
||||
).model_dump(exclude_none=True)
|
||||
litellm_params_json = safe_dumps(litellm_params_dict)
|
||||
del vector_store["litellm_params"]
|
||||
|
||||
_new_vector_store = (
|
||||
await prisma_client.db.litellm_managedvectorstorestable.create(
|
||||
data={
|
||||
**vector_store,
|
||||
"litellm_params": litellm_params_json,
|
||||
}
|
||||
detail="vector_store_id and custom_llm_provider are required"
|
||||
)
|
||||
|
||||
# Extract and validate metadata
|
||||
metadata = vector_store.get("vector_store_metadata")
|
||||
validated_metadata: Optional[Dict] = None
|
||||
if metadata is not None and isinstance(metadata, dict):
|
||||
validated_metadata = metadata
|
||||
|
||||
new_vector_store = await create_vector_store_in_db(
|
||||
vector_store_id=vector_store_id,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
prisma_client=prisma_client,
|
||||
vector_store_name=vector_store.get("vector_store_name"),
|
||||
vector_store_description=vector_store.get("vector_store_description"),
|
||||
vector_store_metadata=validated_metadata,
|
||||
litellm_params=vector_store.get("litellm_params"),
|
||||
created_by=user_api_key_dict.user_id,
|
||||
updated_by=user_api_key_dict.user_id,
|
||||
)
|
||||
|
||||
new_vector_store: LiteLLM_ManagedVectorStore = LiteLLM_ManagedVectorStore(
|
||||
**_new_vector_store.model_dump()
|
||||
)
|
||||
|
||||
# Add vector store to registry
|
||||
if litellm.vector_store_registry is not None:
|
||||
litellm.vector_store_registry.add_vector_store_to_registry(
|
||||
vector_store=new_vector_store
|
||||
)
|
||||
|
||||
return {
|
||||
"status": "success",
|
||||
"message": f"Vector store {vector_store.get('vector_store_id')} created successfully",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue