mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(rag): attach existing OpenAI file ids (#30628)
* fix(rag): attach existing OpenAI file ids * chore: use modern typing in rag ingest fix * chore: retrigger ci
This commit is contained in:
parent
1fda0db66b
commit
64e5685713
7 changed files with 141 additions and 5 deletions
|
|
@ -42,6 +42,8 @@ class BaseRAGIngestion(ABC):
|
|||
vector stores, so it overrides the embedding step to be a no-op.
|
||||
"""
|
||||
|
||||
supports_existing_file_id: bool = False
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
ingest_options: RAGIngestOptions,
|
||||
|
|
@ -280,6 +282,7 @@ class BaseRAGIngestion(ABC):
|
|||
content_type: Optional[str],
|
||||
chunks: List[str],
|
||||
embeddings: Optional[List[List[float]]],
|
||||
existing_file_id: str | None = None,
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
"""
|
||||
Store content in vector store.
|
||||
|
|
@ -292,6 +295,7 @@ class BaseRAGIngestion(ABC):
|
|||
content_type: MIME type
|
||||
chunks: Text chunks (if chunking was done locally)
|
||||
embeddings: Embeddings (if embedding was done locally)
|
||||
existing_file_id: Provider file ID supplied by the caller, if any
|
||||
|
||||
Returns:
|
||||
Tuple of (vector_store_id, file_id)
|
||||
|
|
@ -326,6 +330,12 @@ class BaseRAGIngestion(ABC):
|
|||
)
|
||||
|
||||
try:
|
||||
if existing_file_id and not self.supports_existing_file_id:
|
||||
raise ValueError(
|
||||
f"{self.__class__.__name__} does not support ingesting an existing file_id. "
|
||||
"Upload file data or provide file_url instead."
|
||||
)
|
||||
|
||||
# Step 2: OCR (optional)
|
||||
extracted_text = await self.ocr(
|
||||
file_content=file_content,
|
||||
|
|
@ -349,6 +359,7 @@ class BaseRAGIngestion(ABC):
|
|||
content_type=content_type,
|
||||
chunks=chunks,
|
||||
embeddings=embeddings,
|
||||
existing_file_id=existing_file_id,
|
||||
)
|
||||
|
||||
return RAGIngestResponse(
|
||||
|
|
|
|||
|
|
@ -685,6 +685,7 @@ class BedrockRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
content_type: Optional[str],
|
||||
chunks: List[str],
|
||||
embeddings: Optional[List[List[float]]],
|
||||
existing_file_id: str | None = None,
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
"""
|
||||
Store content in Bedrock Knowledge Base.
|
||||
|
|
@ -701,6 +702,7 @@ class BedrockRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
content_type: MIME type
|
||||
chunks: Ignored - Bedrock handles chunking
|
||||
embeddings: Ignored - Bedrock handles embedding
|
||||
existing_file_id: Existing provider file ID, unsupported for Bedrock
|
||||
|
||||
Returns:
|
||||
Tuple of (knowledge_base_id, file_key)
|
||||
|
|
|
|||
|
|
@ -61,6 +61,7 @@ class GeminiRAGIngestion(BaseRAGIngestion):
|
|||
content_type: Optional[str],
|
||||
chunks: List[str],
|
||||
embeddings: Optional[List[List[float]]],
|
||||
existing_file_id: str | None = None,
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
"""
|
||||
Store content in Gemini File Search store.
|
||||
|
|
@ -75,6 +76,7 @@ class GeminiRAGIngestion(BaseRAGIngestion):
|
|||
content_type: MIME type
|
||||
chunks: Ignored - Gemini handles chunking
|
||||
embeddings: Ignored - Gemini handles embedding
|
||||
existing_file_id: Existing provider file ID, unsupported for Gemini
|
||||
|
||||
Returns:
|
||||
Tuple of (vector_store_id, file_id)
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ so this implementation skips the embedding step and directly uploads files.
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, cast
|
||||
from typing import TYPE_CHECKING, Any, List, Optional, Tuple, cast
|
||||
|
||||
import litellm
|
||||
from litellm.rag.ingestion.base_ingestion import BaseRAGIngestion
|
||||
|
|
@ -29,6 +29,8 @@ class OpenAIRAGIngestion(BaseRAGIngestion):
|
|||
- Chunking is done by OpenAI's vector store (uses 'auto' strategy)
|
||||
"""
|
||||
|
||||
supports_existing_file_id = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
ingest_options: "RAGIngestOptions",
|
||||
|
|
@ -56,6 +58,7 @@ class OpenAIRAGIngestion(BaseRAGIngestion):
|
|||
content_type: Optional[str],
|
||||
chunks: List[str],
|
||||
embeddings: Optional[List[List[float]]],
|
||||
existing_file_id: str | None = None,
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
"""
|
||||
Store content in OpenAI vector store.
|
||||
|
|
@ -71,6 +74,7 @@ class OpenAIRAGIngestion(BaseRAGIngestion):
|
|||
content_type: MIME type
|
||||
chunks: Ignored - OpenAI handles chunking
|
||||
embeddings: Ignored - OpenAI handles embedding
|
||||
existing_file_id: Existing OpenAI file ID to attach
|
||||
|
||||
Returns:
|
||||
Tuple of (vector_store_id, file_id)
|
||||
|
|
@ -82,6 +86,11 @@ class OpenAIRAGIngestion(BaseRAGIngestion):
|
|||
api_key = self.vector_store_config.get("api_key")
|
||||
api_base = self.vector_store_config.get("api_base")
|
||||
|
||||
if existing_file_id and not vector_store_id:
|
||||
raise ValueError(
|
||||
"vector_store_id is required when ingesting an existing file_id"
|
||||
)
|
||||
|
||||
# Create vector store if not provided
|
||||
if not vector_store_id:
|
||||
expires_after = (
|
||||
|
|
@ -96,9 +105,20 @@ class OpenAIRAGIngestion(BaseRAGIngestion):
|
|||
)
|
||||
vector_store_id = create_response.get("id")
|
||||
|
||||
if existing_file_id and vector_store_id:
|
||||
await vector_store_file_acreate(
|
||||
vector_store_id=vector_store_id,
|
||||
file_id=existing_file_id,
|
||||
custom_llm_provider="openai",
|
||||
chunking_strategy=cast(dict[str, Any] | None, self.chunking_strategy),
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
)
|
||||
return vector_store_id, existing_file_id
|
||||
|
||||
# Upload file and attach to vector store
|
||||
result_file_id = None
|
||||
if file_content and filename and vector_store_id:
|
||||
if file_content is not None and filename and vector_store_id:
|
||||
# Upload file to OpenAI
|
||||
file_response = await litellm.acreate_file(
|
||||
file=(
|
||||
|
|
@ -118,9 +138,7 @@ class OpenAIRAGIngestion(BaseRAGIngestion):
|
|||
vector_store_id=vector_store_id,
|
||||
file_id=result_file_id,
|
||||
custom_llm_provider="openai",
|
||||
chunking_strategy=cast(
|
||||
Optional[Dict[str, Any]], self.chunking_strategy
|
||||
),
|
||||
chunking_strategy=cast(dict[str, Any] | None, self.chunking_strategy),
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -464,6 +464,7 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
content_type: Optional[str],
|
||||
chunks: List[str],
|
||||
embeddings: Optional[List[List[float]]],
|
||||
existing_file_id: str | None = None,
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
"""
|
||||
Store vectors in S3 Vectors using PutVectors API.
|
||||
|
|
@ -480,6 +481,7 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
content_type: MIME type (not used for S3 Vectors)
|
||||
chunks: Text chunks
|
||||
embeddings: Vector embeddings
|
||||
existing_file_id: Existing provider file ID, unsupported for S3 Vectors
|
||||
|
||||
Returns:
|
||||
Tuple of (index_name, filename)
|
||||
|
|
|
|||
|
|
@ -74,6 +74,7 @@ class VertexAIRAGIngestion(BaseRAGIngestion, VertexBase):
|
|||
content_type: Optional[str],
|
||||
chunks: List[str],
|
||||
embeddings: Optional[List[List[float]]],
|
||||
existing_file_id: str | None = None,
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
"""
|
||||
Store content in Vertex AI RAG corpus.
|
||||
|
|
@ -88,6 +89,7 @@ class VertexAIRAGIngestion(BaseRAGIngestion, VertexBase):
|
|||
content_type: MIME type
|
||||
chunks: Ignored - Vertex AI handles chunking
|
||||
embeddings: Ignored - Vertex AI handles embedding
|
||||
existing_file_id: Existing provider file ID, unsupported for Vertex AI
|
||||
|
||||
Returns:
|
||||
Tuple of (rag_corpus_id, file_id)
|
||||
|
|
|
|||
99
tests/test_litellm/test_rag_openai_ingestion.py
Normal file
99
tests/test_litellm/test_rag_openai_ingestion.py
Normal file
|
|
@ -0,0 +1,99 @@
|
|||
import asyncio
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
from litellm.rag.ingestion.base_ingestion import BaseRAGIngestion
|
||||
from litellm.rag.ingestion.openai_ingestion import OpenAIRAGIngestion
|
||||
|
||||
|
||||
def test_openai_ingest_existing_file_id_attaches_without_uploading():
|
||||
asyncio.run(_run_openai_existing_file_id_attach_test())
|
||||
|
||||
|
||||
async def _run_openai_existing_file_id_attach_test():
|
||||
ingestion = OpenAIRAGIngestion(
|
||||
{
|
||||
"chunking_strategy": {"type": "auto"},
|
||||
"vector_store": {
|
||||
"custom_llm_provider": "openai",
|
||||
"vector_store_id": "vs_existing",
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.rag.ingestion.openai_ingestion.vector_store_file_acreate",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_attach,
|
||||
patch(
|
||||
"litellm.rag.ingestion.openai_ingestion.litellm.acreate_file",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_upload,
|
||||
):
|
||||
response = await ingestion.ingest(file_id="file_existing")
|
||||
|
||||
assert response["status"] == "completed"
|
||||
assert response["vector_store_id"] == "vs_existing"
|
||||
assert response["file_id"] == "file_existing"
|
||||
mock_upload.assert_not_called()
|
||||
mock_attach.assert_awaited_once_with(
|
||||
vector_store_id="vs_existing",
|
||||
file_id="file_existing",
|
||||
custom_llm_provider="openai",
|
||||
chunking_strategy={"type": "auto"},
|
||||
api_key=None,
|
||||
api_base=None,
|
||||
)
|
||||
|
||||
|
||||
def test_openai_ingest_existing_file_id_requires_vector_store_id():
|
||||
asyncio.run(_run_openai_existing_file_id_requires_vector_store_id_test())
|
||||
|
||||
|
||||
async def _run_openai_existing_file_id_requires_vector_store_id_test():
|
||||
ingestion = OpenAIRAGIngestion({"vector_store": {"custom_llm_provider": "openai"}})
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.rag.ingestion.openai_ingestion.vector_store_acreate",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_create_vector_store,
|
||||
patch(
|
||||
"litellm.rag.ingestion.openai_ingestion.vector_store_file_acreate",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_attach,
|
||||
):
|
||||
response = await ingestion.ingest(file_id="file_existing")
|
||||
|
||||
assert response["status"] == "failed"
|
||||
assert "vector_store_id is required" in response["error"]
|
||||
mock_create_vector_store.assert_not_called()
|
||||
mock_attach.assert_not_called()
|
||||
|
||||
|
||||
class UnsupportedExistingFileIngestion(BaseRAGIngestion):
|
||||
async def store(
|
||||
self,
|
||||
file_content: bytes | None,
|
||||
filename: str | None,
|
||||
content_type: str | None,
|
||||
chunks: list[str],
|
||||
embeddings: list[list[float]] | None,
|
||||
existing_file_id: str | None = None,
|
||||
) -> tuple[str | None, str | None]:
|
||||
raise AssertionError("store should not be called for unsupported file_id")
|
||||
|
||||
|
||||
def test_existing_file_id_fails_for_unsupported_ingestion_provider():
|
||||
asyncio.run(_run_unsupported_existing_file_id_test())
|
||||
|
||||
|
||||
async def _run_unsupported_existing_file_id_test():
|
||||
ingestion = UnsupportedExistingFileIngestion(
|
||||
{"vector_store": {"custom_llm_provider": "unsupported"}}
|
||||
)
|
||||
|
||||
response = await ingestion.ingest(file_id="file_existing")
|
||||
|
||||
assert response["status"] == "failed"
|
||||
assert "does not support ingesting an existing file_id" in response["error"]
|
||||
Loading…
Add table
Reference in a new issue