mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(s3_vectors): embed registered-store ingests with the store's embedding model
The S3 Vectors ingestion embedded every chunk with the request's embedding.model or the default, never the embedding_model the store was registered with, while search on the same store embeds with the registered model. A registered store uploaded to by id alone therefore embedded with the wrong model and AWS rejected the vectors on the dimension mismatch. The store's embedding model now wins for S3 Vectors ingestion through a helper next to the one search already uses
This commit is contained in:
parent
e74a5e0c21
commit
aef209963a
3 changed files with 78 additions and 9 deletions
|
|
@ -8,6 +8,7 @@ from litellm.llms.base_llm.vector_store.transformation import (
|
|||
VectorStoreEmbeddingExecutor,
|
||||
)
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
from litellm.types.rag import RAGIngestEmbeddingOptions
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.vector_stores import (
|
||||
VECTOR_STORE_OPENAI_PARAMS,
|
||||
|
|
@ -57,6 +58,21 @@ def s3_vectors_ingest_target(vector_store_config: Mapping[str, object]) -> tuple
|
|||
return explicit_bucket_name or derived_bucket_name, explicit_index_name or derived_index_name
|
||||
|
||||
|
||||
def s3_vectors_configured_embedding_model(litellm_params: Mapping[str, object]) -> str | None:
|
||||
return _non_empty_str(litellm_params.get("litellm_embedding_model") or litellm_params.get("embedding_model"))
|
||||
|
||||
|
||||
def s3_vectors_ingest_embedding_options(
|
||||
vector_store_config: Mapping[str, object],
|
||||
embedding_options: RAGIngestEmbeddingOptions | None,
|
||||
) -> RAGIngestEmbeddingOptions | None:
|
||||
store_embedding_model: Final = s3_vectors_configured_embedding_model(vector_store_config)
|
||||
if store_embedding_model is None:
|
||||
return embedding_options
|
||||
store_embedding_options: Final[RAGIngestEmbeddingOptions] = {"model": store_embedding_model}
|
||||
return store_embedding_options
|
||||
|
||||
|
||||
class S3VectorsVectorStoreConfig(BaseQueryEmbeddingVectorStoreConfig, BaseAWSLLM):
|
||||
"""Vector store configuration for AWS S3 Vectors."""
|
||||
|
||||
|
|
@ -98,8 +114,7 @@ class S3VectorsVectorStoreConfig(BaseQueryEmbeddingVectorStoreConfig, BaseAWSLLM
|
|||
|
||||
@staticmethod
|
||||
def query_embedding_model(litellm_params: Mapping[str, object]) -> str:
|
||||
configured: Final = litellm_params.get("litellm_embedding_model") or litellm_params.get("embedding_model")
|
||||
return configured if isinstance(configured, str) and configured else _DEFAULT_QUERY_EMBEDDING_MODEL
|
||||
return s3_vectors_configured_embedding_model(litellm_params) or _DEFAULT_QUERY_EMBEDDING_MODEL
|
||||
|
||||
@staticmethod
|
||||
def _query_target(vector_store_id: str, litellm_params: Mapping[str, object]) -> tuple[str, str]:
|
||||
|
|
|
|||
|
|
@ -33,7 +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_ingest_target
|
||||
from litellm.llms.s3_vectors.vector_stores.transformation import (
|
||||
s3_vectors_ingest_embedding_options,
|
||||
s3_vectors_ingest_target,
|
||||
)
|
||||
from litellm.rag.ingestion.base_ingestion import BaseRAGIngestion
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -91,6 +94,7 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
BaseAWSLLM.__init__(self)
|
||||
|
||||
self.vector_bucket_name, self.index_name = s3_vectors_ingest_target(self.vector_store_config)
|
||||
self.embedding_config = s3_vectors_ingest_embedding_options(self.vector_store_config, self.embedding_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",
|
||||
|
|
|
|||
|
|
@ -1,18 +1,68 @@
|
|||
from types import SimpleNamespace
|
||||
|
||||
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'"
|
||||
REQUEST_EMBEDDING_MODEL = "text-embedding-3-small"
|
||||
STORE_EMBEDDING_MODEL = "text-embedding-3-large"
|
||||
REQUEST_EMBEDDING = {"model": REQUEST_EMBEDDING_MODEL}
|
||||
|
||||
|
||||
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},
|
||||
}
|
||||
class _RecordingRouter:
|
||||
def __init__(self):
|
||||
self.embedding_models = []
|
||||
|
||||
async def aembedding(self, model, input):
|
||||
self.embedding_models.append(model)
|
||||
return SimpleNamespace(data=[{"embedding": [0.1, 0.2]} for _ in input])
|
||||
|
||||
|
||||
def _ingestion(embedding=REQUEST_EMBEDDING, router=None, **vector_store):
|
||||
vector_store_options = {"custom_llm_provider": "s3_vectors", "aws_region_name": "us-west-2", **vector_store}
|
||||
ingest_options = {"vector_store": vector_store_options} if embedding is None else {
|
||||
"embedding": embedding,
|
||||
"vector_store": vector_store_options,
|
||||
}
|
||||
return S3VectorsRAGIngestion(ingest_options=ingest_options, router=router)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("store_model_key", ["embedding_model", "litellm_embedding_model"])
|
||||
async def test_a_registered_store_embedding_model_wins_over_the_request_on_ingest(store_model_key):
|
||||
router = _RecordingRouter()
|
||||
ingestion = _ingestion(
|
||||
router=router, vector_store_id="my-embeddings:my-index", **{store_model_key: STORE_EMBEDDING_MODEL}
|
||||
)
|
||||
|
||||
await ingestion.embed(["chunk one", "chunk two"])
|
||||
|
||||
assert router.embedding_models == [STORE_EMBEDDING_MODEL]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_registered_store_embedding_model_is_used_when_the_request_names_none():
|
||||
router = _RecordingRouter()
|
||||
ingestion = _ingestion(
|
||||
embedding=None, router=router, vector_store_id="my-embeddings:my-index", embedding_model=STORE_EMBEDDING_MODEL
|
||||
)
|
||||
|
||||
await ingestion.embed(["chunk"])
|
||||
|
||||
assert router.embedding_models == [STORE_EMBEDDING_MODEL]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("store_model", [{}, {"embedding_model": ""}])
|
||||
async def test_the_request_embedding_model_is_kept_when_the_store_names_none(store_model):
|
||||
router = _RecordingRouter()
|
||||
ingestion = _ingestion(router=router, vector_store_id="my-embeddings:my-index", **store_model)
|
||||
|
||||
await ingestion.embed(["chunk"])
|
||||
|
||||
assert router.embedding_models == [REQUEST_EMBEDDING_MODEL]
|
||||
|
||||
|
||||
def test_store_id_alone_names_the_bucket_and_index():
|
||||
ingestion = _ingestion(vector_store_id="my-embeddings:my-index")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue