mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-05 18:24:49 +00:00
Merge pull request #21120 from BerriAI/litellm_add_rag_ingest_vertex_ai
Add rag ingest vertex ai
This commit is contained in:
@@ -115,8 +115,13 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase):
|
||||
vertex_project = self.get_vertex_ai_project(litellm_params)
|
||||
vertex_location = self.get_vertex_ai_location(litellm_params)
|
||||
|
||||
# Construct full rag corpus path
|
||||
full_rag_corpus = f"projects/{vertex_project}/locations/{vertex_location}/ragCorpora/{vector_store_id}"
|
||||
# Handle both full corpus path and just corpus ID
|
||||
if vector_store_id.startswith("projects/"):
|
||||
# Already a full path
|
||||
full_rag_corpus = vector_store_id
|
||||
else:
|
||||
# Just the corpus ID, construct full path
|
||||
full_rag_corpus = f"projects/{vertex_project}/locations/{vertex_location}/ragCorpora/{vector_store_id}"
|
||||
|
||||
# Build the request body for Vertex AI RAG API
|
||||
request_body: Dict[str, Any] = {
|
||||
|
||||
@@ -4,11 +4,17 @@ RAG Ingestion classes for different providers.
|
||||
|
||||
from litellm.rag.ingestion.base_ingestion import BaseRAGIngestion
|
||||
from litellm.rag.ingestion.bedrock_ingestion import BedrockRAGIngestion
|
||||
from litellm.rag.ingestion.gemini_ingestion import GeminiRAGIngestion
|
||||
from litellm.rag.ingestion.openai_ingestion import OpenAIRAGIngestion
|
||||
from litellm.rag.ingestion.s3_vectors_ingestion import S3VectorsRAGIngestion
|
||||
from litellm.rag.ingestion.vertex_ai_ingestion import VertexAIRAGIngestion
|
||||
|
||||
__all__ = [
|
||||
"BaseRAGIngestion",
|
||||
"BedrockRAGIngestion",
|
||||
"GeminiRAGIngestion",
|
||||
"OpenAIRAGIngestion",
|
||||
"S3VectorsRAGIngestion",
|
||||
"VertexAIRAGIngestion",
|
||||
]
|
||||
|
||||
|
||||
@@ -0,0 +1,478 @@
|
||||
"""
|
||||
Vertex AI-specific RAG Ingestion implementation.
|
||||
|
||||
Vertex AI RAG Engine handles embedding and chunking internally when files are uploaded,
|
||||
so this implementation skips the embedding step and directly uploads files to RAG corpora.
|
||||
|
||||
Based on: https://docs.cloud.google.com/vertex-ai/generative-ai/docs/model-reference/rag-api-v1
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, cast
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.llms.vertex_ai.common_utils import get_vertex_base_url
|
||||
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
|
||||
from litellm.rag.ingestion.base_ingestion import BaseRAGIngestion
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm import Router
|
||||
from litellm.types.rag import RAGIngestOptions
|
||||
|
||||
|
||||
class VertexAIRAGIngestion(BaseRAGIngestion, VertexBase):
|
||||
"""
|
||||
Vertex AI RAG Engine ingestion implementation.
|
||||
|
||||
Key differences from base:
|
||||
- Embedding is handled by Vertex AI RAG Engine when files are uploaded
|
||||
- Files are uploaded using the RAG API (import or upload)
|
||||
- Chunking is done by Vertex AI RAG Engine (supports custom chunking config)
|
||||
- Supports Google Cloud Storage (GCS) and Google Drive sources
|
||||
- Supports custom parsing configurations (layout parser, LLM parser)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
ingest_options: "RAGIngestOptions",
|
||||
router: Optional["Router"] = None,
|
||||
):
|
||||
BaseRAGIngestion.__init__(self, ingest_options=ingest_options, router=router)
|
||||
VertexBase.__init__(self)
|
||||
|
||||
# Extract Vertex AI specific configs from vector_store_config
|
||||
litellm_params = dict(self.vector_store_config)
|
||||
|
||||
# Get project, location, and credentials using VertexBase methods
|
||||
self.project_id = self.safe_get_vertex_ai_project(litellm_params)
|
||||
self.location = self.get_vertex_ai_location(litellm_params) or "us-central1"
|
||||
self.vertex_credentials = self.safe_get_vertex_ai_credentials(litellm_params)
|
||||
|
||||
async def embed(
|
||||
self,
|
||||
chunks: List[str],
|
||||
) -> Optional[List[List[float]]]:
|
||||
"""
|
||||
Vertex AI RAG Engine handles embedding internally - skip this step.
|
||||
|
||||
Returns:
|
||||
None (Vertex AI embeds when files are uploaded to RAG corpus)
|
||||
"""
|
||||
# Vertex AI RAG Engine handles embedding when files are uploaded
|
||||
return None
|
||||
|
||||
async def store(
|
||||
self,
|
||||
file_content: Optional[bytes],
|
||||
filename: Optional[str],
|
||||
content_type: Optional[str],
|
||||
chunks: List[str],
|
||||
embeddings: Optional[List[List[float]]],
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
"""
|
||||
Store content in Vertex AI RAG corpus.
|
||||
|
||||
Vertex AI workflow:
|
||||
1. Create RAG corpus (if not provided)
|
||||
2. Upload file using RAG API (Vertex AI handles chunking/embedding)
|
||||
|
||||
Args:
|
||||
file_content: Raw file bytes
|
||||
filename: Name of the file
|
||||
content_type: MIME type
|
||||
chunks: Ignored - Vertex AI handles chunking
|
||||
embeddings: Ignored - Vertex AI handles embedding
|
||||
|
||||
Returns:
|
||||
Tuple of (rag_corpus_id, file_id)
|
||||
"""
|
||||
if not self.project_id:
|
||||
raise ValueError(
|
||||
"vertex_project is required for Vertex AI RAG ingestion. "
|
||||
"Set it in vector_store config."
|
||||
)
|
||||
|
||||
# Get or create RAG corpus
|
||||
rag_corpus_id = self.vector_store_config.get("vector_store_id")
|
||||
if not rag_corpus_id:
|
||||
rag_corpus_id = await self._create_rag_corpus(
|
||||
display_name=self.ingest_name or "litellm-rag-corpus",
|
||||
description=self.vector_store_config.get("description"),
|
||||
)
|
||||
|
||||
# Upload file to RAG corpus
|
||||
result_file_id = None
|
||||
if file_content and filename and rag_corpus_id:
|
||||
result_file_id = await self._upload_file_to_corpus(
|
||||
rag_corpus_id=rag_corpus_id,
|
||||
filename=filename,
|
||||
file_content=file_content,
|
||||
content_type=content_type,
|
||||
)
|
||||
|
||||
return rag_corpus_id, result_file_id
|
||||
|
||||
async def _create_rag_corpus(
|
||||
self,
|
||||
display_name: str,
|
||||
description: Optional[str] = None,
|
||||
) -> str:
|
||||
"""
|
||||
Create a Vertex AI RAG corpus.
|
||||
|
||||
Args:
|
||||
display_name: Display name for the corpus
|
||||
description: Optional description
|
||||
|
||||
Returns:
|
||||
RAG corpus ID (format: projects/{project}/locations/{location}/ragCorpora/{corpus_id})
|
||||
"""
|
||||
# Get access token using VertexBase method
|
||||
access_token, project_id = self._ensure_access_token(
|
||||
credentials=self.vertex_credentials,
|
||||
project_id=self.project_id,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
|
||||
# Use the project_id from token if not set
|
||||
if not self.project_id:
|
||||
self.project_id = project_id
|
||||
|
||||
# Construct URL using vertex base URL helper
|
||||
base_url = get_vertex_base_url(self.location)
|
||||
url = (
|
||||
f"{base_url}/v1beta1/"
|
||||
f"projects/{self.project_id}/locations/{self.location}/ragCorpora"
|
||||
)
|
||||
|
||||
# Build request body with camelCase keys (Vertex AI API format)
|
||||
request_body: Dict[str, Any] = {
|
||||
"displayName": display_name,
|
||||
}
|
||||
|
||||
if description:
|
||||
request_body["description"] = description
|
||||
|
||||
# Add vector database config if specified
|
||||
vector_db_config = self.vector_store_config.get("vector_db_config")
|
||||
if vector_db_config:
|
||||
request_body["vectorDbConfig"] = vector_db_config
|
||||
|
||||
# Add embedding model config if specified
|
||||
embedding_model = self.vector_store_config.get("embedding_model")
|
||||
if embedding_model:
|
||||
if "vectorDbConfig" not in request_body:
|
||||
request_body["vectorDbConfig"] = {}
|
||||
request_body["vectorDbConfig"]["ragEmbeddingModelConfig"] = {
|
||||
"vertexPredictionEndpoint": {
|
||||
"endpoint": embedding_model
|
||||
}
|
||||
}
|
||||
|
||||
verbose_logger.debug(f"Creating RAG corpus: {url}")
|
||||
verbose_logger.debug(f"Request body: {json.dumps(request_body, indent=2)}")
|
||||
|
||||
client = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.RAG,
|
||||
params={"timeout": 60.0},
|
||||
)
|
||||
|
||||
response = await client.post(
|
||||
url,
|
||||
json=request_body,
|
||||
headers={
|
||||
"Authorization": f"Bearer {access_token}",
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
)
|
||||
if response.status_code not in [200, 201]:
|
||||
error_msg = f"Failed to create RAG corpus: {response.text}"
|
||||
verbose_logger.error(error_msg)
|
||||
raise Exception(error_msg)
|
||||
|
||||
response_data = response.json()
|
||||
verbose_logger.debug(f"Create corpus response: {json.dumps(response_data, indent=2)}")
|
||||
|
||||
# The response is a long-running operation
|
||||
# Check if it's already done or if we need to poll
|
||||
if response_data.get("done"):
|
||||
# Operation completed immediately
|
||||
corpus_name = response_data.get("response", {}).get("name", "")
|
||||
else:
|
||||
# Need to poll the operation
|
||||
operation_name = response_data.get("name", "")
|
||||
verbose_logger.debug(f"Polling operation: {operation_name}")
|
||||
corpus_name = await self._poll_operation(
|
||||
operation_name=operation_name,
|
||||
access_token=access_token,
|
||||
)
|
||||
|
||||
verbose_logger.debug(f"Created RAG corpus: {corpus_name}")
|
||||
return corpus_name
|
||||
|
||||
async def _poll_operation(
|
||||
self,
|
||||
operation_name: str,
|
||||
access_token: str,
|
||||
max_retries: int = 30,
|
||||
retry_delay: float = 2.0,
|
||||
) -> str:
|
||||
"""
|
||||
Poll a long-running operation until it completes.
|
||||
|
||||
Args:
|
||||
operation_name: The operation name (e.g., "operations/123456")
|
||||
access_token: Access token for authentication
|
||||
max_retries: Maximum number of polling attempts
|
||||
retry_delay: Delay between polling attempts in seconds
|
||||
|
||||
Returns:
|
||||
The corpus name from the completed operation
|
||||
|
||||
Raises:
|
||||
Exception: If operation fails or times out
|
||||
"""
|
||||
import asyncio
|
||||
|
||||
base_url = get_vertex_base_url(self.location)
|
||||
# Operation name is like: projects/{project}/locations/{location}/operations/{operation_id}
|
||||
# We need to construct the full URL
|
||||
url = f"{base_url}/v1beta1/{operation_name}"
|
||||
|
||||
client = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.RAG,
|
||||
params={"timeout": 60.0},
|
||||
)
|
||||
|
||||
for attempt in range(max_retries):
|
||||
response = await client.get(
|
||||
url,
|
||||
headers={
|
||||
"Authorization": f"Bearer {access_token}",
|
||||
},
|
||||
)
|
||||
|
||||
if response.status_code != 200:
|
||||
error_msg = f"Failed to poll operation: {response.text}"
|
||||
verbose_logger.error(error_msg)
|
||||
raise Exception(error_msg)
|
||||
|
||||
operation_data = response.json()
|
||||
|
||||
if operation_data.get("done"):
|
||||
# Check for errors
|
||||
if "error" in operation_data:
|
||||
error = operation_data["error"]
|
||||
raise Exception(f"Operation failed: {error}")
|
||||
|
||||
# Extract corpus name from response
|
||||
corpus_name = operation_data.get("response", {}).get("name", "")
|
||||
if corpus_name:
|
||||
return corpus_name
|
||||
else:
|
||||
raise Exception(f"No corpus name in operation response: {operation_data}")
|
||||
|
||||
verbose_logger.debug(f"Operation not done yet, attempt {attempt + 1}/{max_retries}")
|
||||
await asyncio.sleep(retry_delay)
|
||||
|
||||
raise Exception(f"Operation timed out after {max_retries} attempts")
|
||||
|
||||
async def _upload_file_to_corpus(
|
||||
self,
|
||||
rag_corpus_id: str,
|
||||
filename: str,
|
||||
file_content: bytes,
|
||||
content_type: Optional[str],
|
||||
) -> str:
|
||||
"""
|
||||
Upload a file to Vertex AI RAG corpus using multipart upload.
|
||||
|
||||
Args:
|
||||
rag_corpus_id: RAG corpus resource name
|
||||
filename: Name of the file
|
||||
file_content: File content bytes
|
||||
content_type: MIME type
|
||||
|
||||
Returns:
|
||||
File ID or resource name
|
||||
"""
|
||||
# Get access token using VertexBase method
|
||||
access_token, _ = self._ensure_access_token(
|
||||
credentials=self.vertex_credentials,
|
||||
project_id=self.project_id,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
|
||||
# Construct upload URL using vertex base URL helper
|
||||
base_url = get_vertex_base_url(self.location)
|
||||
url = (
|
||||
f"{base_url}/upload/v1beta1/"
|
||||
f"{rag_corpus_id}/ragFiles:upload"
|
||||
)
|
||||
|
||||
# Build metadata for the file with snake_case keys (as per upload API docs)
|
||||
metadata: Dict[str, Any] = {
|
||||
"rag_file": {
|
||||
"display_name": filename,
|
||||
}
|
||||
}
|
||||
|
||||
# Add description if provided
|
||||
description = self.vector_store_config.get("file_description")
|
||||
if description:
|
||||
metadata["rag_file"]["description"] = description
|
||||
|
||||
# Add chunking configuration if provided
|
||||
chunking_strategy = self.chunking_strategy
|
||||
if chunking_strategy and isinstance(chunking_strategy, dict):
|
||||
chunk_size = chunking_strategy.get("chunk_size")
|
||||
chunk_overlap = chunking_strategy.get("chunk_overlap")
|
||||
|
||||
if chunk_size or chunk_overlap:
|
||||
if "upload_rag_file_config" not in metadata:
|
||||
metadata["upload_rag_file_config"] = {}
|
||||
|
||||
metadata["upload_rag_file_config"]["rag_file_transformation_config"] = {
|
||||
"rag_file_chunking_config": {
|
||||
"fixed_length_chunking": {}
|
||||
}
|
||||
}
|
||||
|
||||
chunking_config = metadata["upload_rag_file_config"][
|
||||
"rag_file_transformation_config"
|
||||
]["rag_file_chunking_config"]["fixed_length_chunking"]
|
||||
|
||||
if chunk_size:
|
||||
chunking_config["chunk_size"] = chunk_size
|
||||
if chunk_overlap:
|
||||
chunking_config["chunk_overlap"] = chunk_overlap
|
||||
|
||||
verbose_logger.debug(f"Uploading file to RAG corpus: {url}")
|
||||
verbose_logger.debug(f"Metadata: {json.dumps(metadata, indent=2)}")
|
||||
|
||||
# Prepare multipart form data
|
||||
files = {
|
||||
"metadata": (None, json.dumps(metadata), "application/json"),
|
||||
"file": (filename, file_content, content_type or "application/octet-stream"),
|
||||
}
|
||||
client = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.RAG,
|
||||
params={"timeout": 300.0}, # Longer timeout for large files
|
||||
)
|
||||
|
||||
response = await client.post(
|
||||
url,
|
||||
files=files,
|
||||
headers={
|
||||
"Authorization": f"Bearer {access_token}",
|
||||
"X-Goog-Upload-Protocol": "multipart",
|
||||
},
|
||||
)
|
||||
|
||||
if response.status_code not in [200, 201]:
|
||||
error_msg = f"Failed to upload file: {response.text}"
|
||||
verbose_logger.error(error_msg)
|
||||
raise Exception(error_msg)
|
||||
|
||||
# Parse response to get file ID
|
||||
try:
|
||||
response_data = response.json()
|
||||
# The response should contain the rag_file resource name
|
||||
file_id = response_data.get("ragFile", {}).get("name", "")
|
||||
if not file_id:
|
||||
file_id = response_data.get("name", "")
|
||||
|
||||
verbose_logger.debug(f"Upload complete. File ID: {file_id}")
|
||||
return file_id
|
||||
except Exception as e:
|
||||
verbose_logger.warning(f"Could not parse upload response: {e}")
|
||||
return "uploaded"
|
||||
|
||||
async def _import_files_from_gcs(
|
||||
self,
|
||||
rag_corpus_id: str,
|
||||
gcs_uris: List[str],
|
||||
) -> str:
|
||||
"""
|
||||
Import files from Google Cloud Storage into RAG corpus.
|
||||
|
||||
Args:
|
||||
rag_corpus_id: RAG corpus resource name
|
||||
gcs_uris: List of GCS URIs (e.g., ["gs://bucket/file.pdf"])
|
||||
|
||||
Returns:
|
||||
Operation name for tracking import progress
|
||||
"""
|
||||
# Get access token using VertexBase method
|
||||
access_token, _ = self._ensure_access_token(
|
||||
credentials=self.vertex_credentials,
|
||||
project_id=self.project_id,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
|
||||
# Construct import URL using vertex base URL helper
|
||||
base_url = get_vertex_base_url(self.location)
|
||||
url = (
|
||||
f"{base_url}/v1beta1/"
|
||||
f"{rag_corpus_id}/ragFiles:import"
|
||||
)
|
||||
|
||||
# Build request body with camelCase keys (Vertex AI API format)
|
||||
request_body: Dict[str, Any] = {
|
||||
"importRagFilesConfig": {
|
||||
"gcsSource": {
|
||||
"uris": gcs_uris
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
# Add chunking configuration if provided
|
||||
chunking_strategy = self.chunking_strategy
|
||||
if chunking_strategy and isinstance(chunking_strategy, dict):
|
||||
chunk_size = chunking_strategy.get("chunk_size")
|
||||
chunk_overlap = chunking_strategy.get("chunk_overlap")
|
||||
|
||||
if chunk_size or chunk_overlap:
|
||||
request_body["importRagFilesConfig"]["ragFileChunkingConfig"] = {
|
||||
"chunkSize": chunk_size or 1024,
|
||||
"chunkOverlap": chunk_overlap or 200,
|
||||
}
|
||||
|
||||
# Add max embedding requests per minute if specified
|
||||
max_embedding_qpm = self.vector_store_config.get("max_embedding_requests_per_min")
|
||||
if max_embedding_qpm:
|
||||
request_body["importRagFilesConfig"]["maxEmbeddingRequestsPerMin"] = max_embedding_qpm
|
||||
|
||||
verbose_logger.debug(f"Importing files from GCS: {url}")
|
||||
verbose_logger.debug(f"Request body: {json.dumps(request_body, indent=2)}")
|
||||
|
||||
client = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.RAG,
|
||||
params={"timeout": 60.0},
|
||||
)
|
||||
|
||||
response = await client.post(
|
||||
url,
|
||||
json=request_body,
|
||||
headers={
|
||||
"Authorization": f"Bearer {access_token}",
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
)
|
||||
|
||||
if response.status_code not in [200, 201]:
|
||||
error_msg = f"Failed to import files: {response.text}"
|
||||
verbose_logger.error(error_msg)
|
||||
raise Exception(error_msg)
|
||||
|
||||
response_data = response.json()
|
||||
operation_name = response_data.get("name", "")
|
||||
|
||||
verbose_logger.debug(f"Import operation started: {operation_name}")
|
||||
return operation_name
|
||||
@@ -32,6 +32,7 @@ from litellm.rag.ingestion.bedrock_ingestion import BedrockRAGIngestion
|
||||
from litellm.rag.ingestion.gemini_ingestion import GeminiRAGIngestion
|
||||
from litellm.rag.ingestion.openai_ingestion import OpenAIRAGIngestion
|
||||
from litellm.rag.ingestion.s3_vectors_ingestion import S3VectorsRAGIngestion
|
||||
from litellm.rag.ingestion.vertex_ai_ingestion import VertexAIRAGIngestion
|
||||
from litellm.rag.rag_query import RAGQuery
|
||||
from litellm.types.rag import (
|
||||
RAGIngestOptions,
|
||||
@@ -50,6 +51,7 @@ INGESTION_REGISTRY: Dict[str, Type[BaseRAGIngestion]] = {
|
||||
"bedrock": BedrockRAGIngestion,
|
||||
"gemini": GeminiRAGIngestion,
|
||||
"s3_vectors": S3VectorsRAGIngestion,
|
||||
"vertex_ai": VertexAIRAGIngestion,
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -1,14 +1,19 @@
|
||||
"""
|
||||
Vertex AI RAG Engine ingestion tests.
|
||||
|
||||
Tests the Vertex AI RAG ingestion implementation that:
|
||||
- Creates RAG corpora automatically (or uses existing ones)
|
||||
- Uploads files directly to Vertex AI RAG Engine
|
||||
- Handles long-running operations for corpus creation
|
||||
- Supports both file upload and GCS import
|
||||
|
||||
Requires:
|
||||
- gcloud auth application-default login (for ADC authentication)
|
||||
|
||||
Environment variables:
|
||||
- VERTEX_PROJECT: GCP project ID (required)
|
||||
- VERTEX_LOCATION: GCP region (optional, defaults to europe-west1)
|
||||
- VERTEX_CORPUS_ID: Existing RAG corpus ID (required for Vertex AI)
|
||||
- GCS_BUCKET_NAME: GCS bucket for file uploads (required)
|
||||
- VERTEX_LOCATION: GCP region (optional, defaults to us-central1)
|
||||
- VERTEX_CORPUS_ID: Existing RAG corpus ID (optional - will create if not provided)
|
||||
"""
|
||||
|
||||
import os
|
||||
@@ -31,37 +36,24 @@ class TestRAGVertexAI(BaseRAGTest):
|
||||
def check_env_vars(self):
|
||||
"""Check required environment variables before each test."""
|
||||
vertex_project = os.environ.get("VERTEX_PROJECT")
|
||||
corpus_id = os.environ.get("VERTEX_CORPUS_ID")
|
||||
gcs_bucket = os.environ.get("GCS_BUCKET_NAME")
|
||||
|
||||
if not vertex_project:
|
||||
pytest.skip("Skipping Vertex AI test: VERTEX_PROJECT required")
|
||||
|
||||
if not corpus_id:
|
||||
pytest.skip("Skipping Vertex AI test: VERTEX_CORPUS_ID required")
|
||||
|
||||
if not gcs_bucket:
|
||||
pytest.skip("Skipping Vertex AI test: GCS_BUCKET_NAME required")
|
||||
|
||||
# Check if vertexai is installed
|
||||
try:
|
||||
from vertexai import rag
|
||||
except ImportError:
|
||||
pytest.skip("Skipping Vertex AI test: google-cloud-aiplatform>=1.60.0 required")
|
||||
|
||||
def get_base_ingest_options(self) -> RAGIngestOptions:
|
||||
"""
|
||||
Return Vertex AI-specific ingest options.
|
||||
|
||||
Chunking is configured via chunking_strategy (unified interface),
|
||||
not inside vector_store.
|
||||
|
||||
If VERTEX_CORPUS_ID is not set, a new corpus will be created automatically.
|
||||
"""
|
||||
corpus_id = os.environ.get("VERTEX_CORPUS_ID")
|
||||
vertex_project = os.environ.get("VERTEX_PROJECT")
|
||||
vertex_location = os.environ.get("VERTEX_LOCATION", "europe-west1")
|
||||
gcs_bucket = os.environ.get("GCS_BUCKET_NAME")
|
||||
vertex_location = os.environ.get("VERTEX_LOCATION", "us-central1")
|
||||
corpus_id = os.environ.get("VERTEX_CORPUS_ID") # Optional
|
||||
|
||||
return {
|
||||
options: RAGIngestOptions = {
|
||||
"chunking_strategy": {
|
||||
"chunk_size": 512,
|
||||
"chunk_overlap": 100,
|
||||
@@ -70,61 +62,174 @@ class TestRAGVertexAI(BaseRAGTest):
|
||||
"custom_llm_provider": "vertex_ai",
|
||||
"vertex_project": vertex_project,
|
||||
"vertex_location": vertex_location,
|
||||
"vector_store_id": corpus_id,
|
||||
"gcs_bucket": gcs_bucket,
|
||||
"wait_for_import": True,
|
||||
},
|
||||
}
|
||||
|
||||
# Add corpus ID if provided (otherwise will create new corpus)
|
||||
if corpus_id:
|
||||
options["vector_store"]["vector_store_id"] = corpus_id
|
||||
|
||||
return options
|
||||
|
||||
async def query_vector_store(
|
||||
self,
|
||||
vector_store_id: str,
|
||||
query: str,
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""Query Vertex AI RAG corpus."""
|
||||
try:
|
||||
from vertexai import init as vertexai_init
|
||||
from vertexai import rag
|
||||
except ImportError:
|
||||
pytest.skip("vertexai required for Vertex AI tests")
|
||||
|
||||
"""
|
||||
Query Vertex AI RAG corpus using LiteLLM's vector store search.
|
||||
|
||||
Args:
|
||||
vector_store_id: The RAG corpus ID (can be full path or just the ID)
|
||||
query: The search query
|
||||
|
||||
Returns:
|
||||
Search results dict or None if no results found
|
||||
"""
|
||||
vertex_project = os.environ.get("VERTEX_PROJECT")
|
||||
vertex_location = os.environ.get("VERTEX_LOCATION", "europe-west1")
|
||||
vertex_location = os.environ.get("VERTEX_LOCATION", "us-central1")
|
||||
|
||||
# Initialize Vertex AI
|
||||
vertexai_init(project=vertex_project, location=vertex_location)
|
||||
try:
|
||||
# Use LiteLLM's vector store search
|
||||
search_response = await litellm.vector_stores.asearch(
|
||||
vector_store_id=vector_store_id,
|
||||
query=query,
|
||||
max_num_results=5,
|
||||
custom_llm_provider="vertex_ai",
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
)
|
||||
|
||||
# Build corpus name
|
||||
corpus_name = f"projects/{vertex_project}/locations/{vertex_location}/ragCorpora/{vector_store_id}"
|
||||
# Check if we got results
|
||||
if search_response and search_response.get("data"):
|
||||
results = []
|
||||
for item in search_response["data"]:
|
||||
# Extract text from content
|
||||
text = ""
|
||||
if item.get("content"):
|
||||
for content_item in item["content"]:
|
||||
if content_item.get("text"):
|
||||
text += content_item["text"]
|
||||
|
||||
results.append({
|
||||
"text": text,
|
||||
"score": item.get("score", 0.0),
|
||||
"file_id": item.get("file_id", ""),
|
||||
"filename": item.get("filename", ""),
|
||||
})
|
||||
|
||||
# Query the corpus
|
||||
response = rag.retrieval_query(
|
||||
rag_resources=[
|
||||
rag.RagResource(rag_corpus=corpus_name)
|
||||
],
|
||||
text=query,
|
||||
rag_retrieval_config=rag.RagRetrievalConfig(
|
||||
top_k=5,
|
||||
),
|
||||
)
|
||||
# Check if query terms appear in results
|
||||
for result in results:
|
||||
if query.lower() in result["text"].lower():
|
||||
return {"results": results}
|
||||
|
||||
if hasattr(response, 'contexts') and response.contexts.contexts:
|
||||
# Convert to dict format
|
||||
results = []
|
||||
for ctx in response.contexts.contexts:
|
||||
results.append({
|
||||
"text": ctx.text,
|
||||
"score": ctx.score,
|
||||
"source_uri": ctx.source_uri,
|
||||
})
|
||||
# Return results even if exact match not found
|
||||
return {"results": results}
|
||||
|
||||
# Check if query terms appear in results
|
||||
for result in results:
|
||||
if query.lower() in result["text"].lower():
|
||||
return {"results": results}
|
||||
return None
|
||||
|
||||
# Return results even if exact match not found
|
||||
return {"results": results}
|
||||
except Exception as e:
|
||||
print(f"Query failed: {e}")
|
||||
return None
|
||||
|
||||
return None
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_corpus_and_ingest(self):
|
||||
"""
|
||||
Test creating a new RAG corpus and ingesting a file.
|
||||
|
||||
This test specifically validates:
|
||||
- Automatic corpus creation when vector_store_id is not provided
|
||||
- Long-running operation polling for corpus creation
|
||||
- File upload to the newly created corpus
|
||||
"""
|
||||
litellm._turn_on_debug()
|
||||
|
||||
filename, unique_id = self.get_unique_filename("create_corpus")
|
||||
text_content = f"""
|
||||
Test document {unique_id} for Vertex AI RAG corpus creation.
|
||||
This tests the automatic corpus creation feature.
|
||||
The corpus should be created and the file should be uploaded successfully.
|
||||
""".encode("utf-8")
|
||||
file_data = (filename, text_content, "text/plain")
|
||||
|
||||
# Get base options WITHOUT corpus_id to trigger creation
|
||||
ingest_options = self.get_base_ingest_options()
|
||||
# Remove corpus_id if it was set from env var
|
||||
if "vector_store_id" in ingest_options.get("vector_store", {}):
|
||||
del ingest_options["vector_store"]["vector_store_id"]
|
||||
|
||||
ingest_options["name"] = f"test-create-corpus-{unique_id}"
|
||||
|
||||
try:
|
||||
response = await litellm.rag.aingest(
|
||||
ingest_options=ingest_options,
|
||||
file_data=file_data,
|
||||
)
|
||||
|
||||
print(f"Create Corpus Response: {response}")
|
||||
|
||||
# Validate response
|
||||
assert "id" in response
|
||||
assert response["id"].startswith("ingest_")
|
||||
assert "status" in response
|
||||
assert response["status"] == "completed", f"Expected completed, got {response['status']}"
|
||||
assert "vector_store_id" in response
|
||||
assert response["vector_store_id"], "vector_store_id should not be empty"
|
||||
|
||||
# The vector_store_id should be a full corpus path
|
||||
corpus_id = response["vector_store_id"]
|
||||
assert "projects/" in corpus_id, "Corpus ID should be a full resource path"
|
||||
assert "ragCorpora/" in corpus_id, "Corpus ID should contain ragCorpora"
|
||||
|
||||
print(f"✓ Successfully created corpus: {corpus_id}")
|
||||
print(f"✓ Successfully uploaded file: {response.get('file_id')}")
|
||||
|
||||
except litellm.InternalServerError as e:
|
||||
pytest.skip(f"Skipping test due to litellm.InternalServerError: {e}")
|
||||
except Exception as e:
|
||||
print(f"Test failed with error: {e}")
|
||||
raise
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ingest_with_existing_corpus(self):
|
||||
"""
|
||||
Test ingesting a file to an existing RAG corpus.
|
||||
|
||||
This test validates:
|
||||
- Using an existing corpus_id from environment variable
|
||||
- Direct file upload without corpus creation
|
||||
"""
|
||||
corpus_id = os.environ.get("VERTEX_CORPUS_ID")
|
||||
if not corpus_id:
|
||||
pytest.skip("Skipping test: VERTEX_CORPUS_ID not set")
|
||||
|
||||
litellm._turn_on_debug()
|
||||
|
||||
filename, unique_id = self.get_unique_filename("existing_corpus")
|
||||
text_content = f"""
|
||||
Test document {unique_id} for existing Vertex AI RAG corpus.
|
||||
This tests file upload to a pre-existing corpus.
|
||||
""".encode("utf-8")
|
||||
file_data = (filename, text_content, "text/plain")
|
||||
|
||||
ingest_options = self.get_base_ingest_options()
|
||||
ingest_options["name"] = f"test-existing-corpus-{unique_id}"
|
||||
|
||||
try:
|
||||
response = await litellm.rag.aingest(
|
||||
ingest_options=ingest_options,
|
||||
file_data=file_data,
|
||||
)
|
||||
|
||||
print(f"Existing Corpus Ingest Response: {response}")
|
||||
|
||||
assert response["status"] == "completed"
|
||||
assert response["vector_store_id"] == corpus_id or corpus_id in response["vector_store_id"]
|
||||
assert response.get("file_id"), "file_id should be present"
|
||||
|
||||
print(f"✓ Successfully uploaded to existing corpus: {corpus_id}")
|
||||
print(f"✓ File ID: {response.get('file_id')}")
|
||||
|
||||
except litellm.InternalServerError as e:
|
||||
pytest.skip(f"Skipping test due to litellm.InternalServerError: {e}")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user