mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(rag): read a registered S3 Vectors store's bucket and index from its id
A registered S3 Vectors store usually carries only its "bucket:index" id, and the previous commit stopped forwarding the caller's bucket and index for a managed store, so ingesting into one raised KeyError 'vector_bucket_name'. The ingestion now derives both from vector_store_id with the rule the search side already uses, explicit keys still winning. The caller's litellm_credential_name is dropped for a managed store too, since it expands into api_key and api_base, and max_embedding_requests_per_min joins the per-upload options a caller may still set.
This commit is contained in:
parent
a98c48f933
commit
e2d118aaf8
6 changed files with 157 additions and 21 deletions
|
|
@ -26,6 +26,19 @@ else:
|
|||
|
||||
_DEFAULT_QUERY_EMBEDDING_MODEL: Final = "text-embedding-3-small"
|
||||
_DEFAULT_TOP_K: Final = 5
|
||||
S3_VECTORS_STORE_ID_ERROR: Final = (
|
||||
"vector_store_id must be in format 'bucket_name:index_name' for S3 Vectors, "
|
||||
"or vector_bucket_name must be provided in litellm_params"
|
||||
)
|
||||
|
||||
|
||||
def split_s3_vectors_store_id(vector_store_id: str, fallback_bucket_name: object) -> tuple[str, str]:
|
||||
if ":" in vector_store_id:
|
||||
bucket_name, index_name = vector_store_id.split(":", 1)
|
||||
return bucket_name, index_name
|
||||
if not isinstance(fallback_bucket_name, str) or not fallback_bucket_name:
|
||||
raise ValueError(S3_VECTORS_STORE_ID_ERROR)
|
||||
return fallback_bucket_name, vector_store_id
|
||||
|
||||
|
||||
class S3VectorsVectorStoreConfig(BaseQueryEmbeddingVectorStoreConfig, BaseAWSLLM):
|
||||
|
|
@ -74,16 +87,7 @@ class S3VectorsVectorStoreConfig(BaseQueryEmbeddingVectorStoreConfig, BaseAWSLLM
|
|||
|
||||
@staticmethod
|
||||
def _query_target(vector_store_id: str, litellm_params: Mapping[str, object]) -> tuple[str, str]:
|
||||
if ":" in vector_store_id:
|
||||
bucket_name, index_name = vector_store_id.split(":", 1)
|
||||
return bucket_name, index_name
|
||||
bucket_name_from_params: Final = litellm_params.get("vector_bucket_name")
|
||||
if not isinstance(bucket_name_from_params, str) or not bucket_name_from_params:
|
||||
raise ValueError(
|
||||
"vector_store_id must be in format 'bucket_name:index_name' for S3 Vectors, "
|
||||
"or vector_bucket_name must be provided in litellm_params"
|
||||
)
|
||||
return bucket_name_from_params, vector_store_id
|
||||
return split_s3_vectors_store_id(vector_store_id, litellm_params.get("vector_bucket_name"))
|
||||
|
||||
@staticmethod
|
||||
def _query_request(
|
||||
|
|
|
|||
|
|
@ -169,12 +169,12 @@ def _ingest_provider_error(vector_store_config: Mapping[str, object]) -> str | N
|
|||
_MANAGED_STORE_CALLER_OPTIONS: Final = frozenset(
|
||||
{
|
||||
"vector_store_id",
|
||||
"litellm_credential_name",
|
||||
"data_source_id",
|
||||
"wait_for_ingestion",
|
||||
"ingestion_timeout",
|
||||
"custom_metadata",
|
||||
"file_description",
|
||||
"max_embedding_requests_per_min",
|
||||
}
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -33,6 +33,10 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.llms.s3_vectors.vector_stores.transformation import (
|
||||
S3_VECTORS_STORE_ID_ERROR,
|
||||
split_s3_vectors_store_id,
|
||||
)
|
||||
from litellm.rag.ingestion.base_ingestion import BaseRAGIngestion
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -62,6 +66,22 @@ class S3VectorsQueryResponse(TypedDict, total=False):
|
|||
vectors: Sequence[S3VectorsQueryMatch]
|
||||
|
||||
|
||||
def _non_empty_str(value: object) -> str | None:
|
||||
return value if isinstance(value, str) and value else None
|
||||
|
||||
|
||||
def s3_vectors_ingest_target(vector_store_config: Mapping[str, object]) -> tuple[str, str | None]:
|
||||
explicit_bucket_name: Final = _non_empty_str(vector_store_config.get("vector_bucket_name"))
|
||||
explicit_index_name: Final = _non_empty_str(vector_store_config.get("index_name"))
|
||||
vector_store_id: Final = _non_empty_str(vector_store_config.get("vector_store_id"))
|
||||
if vector_store_id is None:
|
||||
if explicit_bucket_name is None:
|
||||
raise ValueError(S3_VECTORS_STORE_ID_ERROR)
|
||||
return explicit_bucket_name, explicit_index_name
|
||||
derived_bucket_name, derived_index_name = split_s3_vectors_store_id(vector_store_id, explicit_bucket_name)
|
||||
return explicit_bucket_name or derived_bucket_name, explicit_index_name or derived_index_name
|
||||
|
||||
|
||||
class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
||||
"""
|
||||
S3 Vectors RAG ingestion using httpx + AWS SigV4 signing.
|
||||
|
|
@ -73,8 +93,9 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
4. Store vectors with PutVectors API
|
||||
|
||||
Configuration:
|
||||
- vector_bucket_name: S3 vector bucket name (required)
|
||||
- index_name: Vector index name (auto-creates if not provided)
|
||||
- vector_store_id: "bucket_name:index_name" of an existing index, or an index name when vector_bucket_name is set
|
||||
- vector_bucket_name: S3 vector bucket name (required unless vector_store_id carries it)
|
||||
- index_name: Vector index name (auto-creates if neither it nor vector_store_id is provided)
|
||||
- dimension: Vector dimension (default: S3_VECTORS_DEFAULT_DIMENSION)
|
||||
- distance_metric: "cosine" or "euclidean" (default: S3_VECTORS_DEFAULT_DISTANCE_METRIC)
|
||||
- non_filterable_metadata_keys: List of metadata keys to exclude from filtering
|
||||
|
|
@ -88,9 +109,7 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
BaseRAGIngestion.__init__(self, ingest_options=ingest_options, router=router)
|
||||
BaseAWSLLM.__init__(self)
|
||||
|
||||
# Extract config
|
||||
self.vector_bucket_name: str = self.vector_store_config["vector_bucket_name"]
|
||||
self.index_name: str | None = self.vector_store_config.get("index_name")
|
||||
self.vector_bucket_name, self.index_name = s3_vectors_ingest_target(self.vector_store_config)
|
||||
self.distance_metric: str = self.vector_store_config.get("distance_metric", S3_VECTORS_DEFAULT_DISTANCE_METRIC)
|
||||
self.non_filterable_metadata_keys: Sequence[str] = self.vector_store_config.get(
|
||||
"non_filterable_metadata_keys",
|
||||
|
|
|
|||
|
|
@ -269,6 +269,17 @@ BEDROCK_REGISTRY_STORE = {
|
|||
"aws_secret_access_key": "registry-secret",
|
||||
},
|
||||
}
|
||||
CREDENTIALED_REGISTRY_STORE = {
|
||||
"vector_store_id": "cred-store",
|
||||
"custom_llm_provider": "openai",
|
||||
"litellm_credential_name": "registry-openai",
|
||||
"litellm_params": {},
|
||||
}
|
||||
VERTEX_REGISTRY_STORE = {
|
||||
"vector_store_id": "projects/registry-project/locations/us-central1/ragCorpora/42",
|
||||
"custom_llm_provider": "vertex_ai",
|
||||
"litellm_params": {"vertex_project": "registry-project", "vertex_location": "us-central1"},
|
||||
}
|
||||
UNSUPPORTED_INGEST_PROVIDER_ERROR = (
|
||||
"Provider '{provider}' is not supported for RAG ingestion. "
|
||||
"Supported providers: openai, bedrock, gemini, s3_vectors, vertex_ai"
|
||||
|
|
@ -426,11 +437,12 @@ def test_rag_ingest_unmanaged_store_keeps_the_callers_full_config(client_interna
|
|||
assert mock_aingest.await_args.kwargs["ingest_options"]["vector_store"] == caller_config
|
||||
|
||||
|
||||
def test_rag_ingest_db_managed_store_keeps_the_callers_credential_name(client_internal_user):
|
||||
def test_rag_ingest_db_managed_store_drops_the_callers_credential_name(client_internal_user):
|
||||
"""
|
||||
A store synced from the database carries litellm_credential_name=None; that
|
||||
null is the absence of a store-side value, not an override, so the credential
|
||||
the caller named must survive the merge exactly as it did before the fix.
|
||||
litellm_credential_name expands into api_key and api_base at ingest time, so a
|
||||
caller naming one would point a managed store's upload at a different endpoint.
|
||||
A store synced from the database carries litellm_credential_name=None, and that
|
||||
null must not resurrect the caller's choice either.
|
||||
"""
|
||||
aingest_patch, registry_patch = _patched_ingest_boundary(
|
||||
DB_MANAGED_STORE, {"vector_store_id": "db-store", "file_id": "file_123"}
|
||||
|
|
@ -447,11 +459,60 @@ def test_rag_ingest_db_managed_store_keeps_the_callers_credential_name(client_in
|
|||
|
||||
assert response.status_code == 200, response.json()
|
||||
forwarded = mock_aingest.await_args.kwargs["ingest_options"]["vector_store"]
|
||||
assert forwarded["litellm_credential_name"] == "team-openai"
|
||||
assert "litellm_credential_name" not in forwarded
|
||||
assert forwarded["custom_llm_provider"] == "openai"
|
||||
assert forwarded["ttl_days"] == 7
|
||||
|
||||
|
||||
def test_rag_ingest_registry_store_credential_name_beats_the_callers(client_internal_user):
|
||||
aingest_patch, registry_patch = _patched_ingest_boundary(
|
||||
CREDENTIALED_REGISTRY_STORE, {"vector_store_id": "cred-store", "file_id": "file_123"}
|
||||
)
|
||||
with (
|
||||
aingest_patch as mock_aingest,
|
||||
registry_patch,
|
||||
_patched_prisma_client(None),
|
||||
):
|
||||
response = client_internal_user.post(
|
||||
"/v1/rag/ingest",
|
||||
**_ingest_form({"vector_store_id": "cred-store", "litellm_credential_name": "team-openai"}),
|
||||
)
|
||||
|
||||
assert response.status_code == 200, response.json()
|
||||
forwarded = mock_aingest.await_args.kwargs["ingest_options"]["vector_store"]
|
||||
assert forwarded["litellm_credential_name"] == "registry-openai"
|
||||
|
||||
|
||||
def test_rag_ingest_registry_store_keeps_the_callers_vertex_embedding_throttle(client_internal_user):
|
||||
aingest_patch, registry_patch = _patched_ingest_boundary(
|
||||
VERTEX_REGISTRY_STORE, {"vector_store_id": VERTEX_REGISTRY_STORE["vector_store_id"], "file_id": "file_123"}
|
||||
)
|
||||
with (
|
||||
aingest_patch as mock_aingest,
|
||||
registry_patch,
|
||||
_patched_prisma_client(None),
|
||||
):
|
||||
response = client_internal_user.post(
|
||||
"/v1/rag/ingest",
|
||||
**_ingest_form(
|
||||
{
|
||||
"vector_store_id": VERTEX_REGISTRY_STORE["vector_store_id"],
|
||||
"max_embedding_requests_per_min": 500,
|
||||
"vector_db_config": {"pinecone": {"index_name": "attacker-index"}},
|
||||
}
|
||||
),
|
||||
)
|
||||
|
||||
assert response.status_code == 200, response.json()
|
||||
assert mock_aingest.await_args.kwargs["ingest_options"]["vector_store"] == {
|
||||
"vector_store_id": VERTEX_REGISTRY_STORE["vector_store_id"],
|
||||
"custom_llm_provider": "vertex_ai",
|
||||
"vertex_project": "registry-project",
|
||||
"vertex_location": "us-central1",
|
||||
"max_embedding_requests_per_min": 500,
|
||||
}
|
||||
|
||||
|
||||
def test_rag_ingest_rejects_registry_store_provider_without_ingestion_support(client_internal_user):
|
||||
"""
|
||||
Regression for LIT-7956: a registry store on a provider with no ingestion
|
||||
|
|
|
|||
0
tests/test_litellm/rag/ingestion/__init__.py
Normal file
0
tests/test_litellm/rag/ingestion/__init__.py
Normal file
|
|
@ -0,0 +1,52 @@
|
|||
import pytest
|
||||
|
||||
from litellm.rag.ingestion.s3_vectors_ingestion import S3VectorsRAGIngestion
|
||||
|
||||
STORE_ID_FORMAT_ERROR = "vector_store_id must be in format 'bucket_name:index_name'"
|
||||
|
||||
|
||||
def _ingestion(**vector_store):
|
||||
return S3VectorsRAGIngestion(
|
||||
ingest_options={
|
||||
"embedding": {"model": "text-embedding-3-small"},
|
||||
"vector_store": {"custom_llm_provider": "s3_vectors", "aws_region_name": "us-west-2", **vector_store},
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def test_store_id_alone_names_the_bucket_and_index():
|
||||
"""
|
||||
Regression for LIT-7956: a registered S3 Vectors store carries only its
|
||||
"bucket:index" id, and the proxy no longer forwards the caller's bucket and
|
||||
index for a managed store, so the ingestion must read both from the id.
|
||||
"""
|
||||
ingestion = _ingestion(vector_store_id="my-embeddings:my-index")
|
||||
|
||||
assert (ingestion.vector_bucket_name, ingestion.index_name) == ("my-embeddings", "my-index")
|
||||
|
||||
|
||||
def test_store_id_without_a_colon_is_the_index_inside_the_given_bucket():
|
||||
ingestion = _ingestion(vector_store_id="my-index", vector_bucket_name="my-embeddings")
|
||||
|
||||
assert (ingestion.vector_bucket_name, ingestion.index_name) == ("my-embeddings", "my-index")
|
||||
|
||||
|
||||
def test_explicit_bucket_and_index_win_over_the_store_id():
|
||||
ingestion = _ingestion(vector_store_id="id-bucket:id-index", vector_bucket_name="my-bucket", index_name="docs")
|
||||
|
||||
assert (ingestion.vector_bucket_name, ingestion.index_name) == ("my-bucket", "docs")
|
||||
|
||||
|
||||
def test_bucket_alone_leaves_the_index_to_be_generated():
|
||||
ingestion = _ingestion(vector_bucket_name="my-embeddings")
|
||||
|
||||
assert (ingestion.vector_bucket_name, ingestion.index_name) == ("my-embeddings", None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"vector_store",
|
||||
[{}, {"vector_store_id": "my-index"}, {"vector_store_id": "my-index", "vector_bucket_name": ""}],
|
||||
)
|
||||
def test_no_bucket_anywhere_is_rejected(vector_store):
|
||||
with pytest.raises(ValueError, match=STORE_ID_FORMAT_ERROR):
|
||||
_ingestion(**vector_store)
|
||||
Loading…
Add table
Reference in a new issue