mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(s3-vectors): use selected embedding model during ingest
This commit is contained in:
parent
73e9071311
commit
44e96a7c1d
2 changed files with 57 additions and 0 deletions
|
|
@ -67,6 +67,9 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
|
||||
# Extract config
|
||||
self.vector_bucket_name = self.vector_store_config["vector_bucket_name"]
|
||||
embedding_model = self.vector_store_config.get("embedding_model")
|
||||
if not self.embedding_config and embedding_model:
|
||||
self.embedding_config = {"model": embedding_model}
|
||||
self.index_name = self.vector_store_config.get("index_name")
|
||||
self.distance_metric = self.vector_store_config.get(
|
||||
"distance_metric", S3_VECTORS_DEFAULT_DISTANCE_METRIC
|
||||
|
|
|
|||
|
|
@ -0,0 +1,54 @@
|
|||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.rag.ingestion.s3_vectors_ingestion import S3VectorsRAGIngestion
|
||||
|
||||
|
||||
def _s3_vectors_ingest_options(**overrides):
|
||||
vector_store = {
|
||||
"custom_llm_provider": "s3_vectors",
|
||||
"vector_bucket_name": "test-vector-bucket",
|
||||
"aws_region_name": "us-east-1",
|
||||
}
|
||||
vector_store.update(overrides.pop("vector_store", {}))
|
||||
return {
|
||||
"vector_store": vector_store,
|
||||
**overrides,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_s3_vectors_ingestion_uses_vector_store_embedding_model():
|
||||
ingestion = S3VectorsRAGIngestion(
|
||||
ingest_options=_s3_vectors_ingest_options(
|
||||
vector_store={"embedding_model": "text-embedding-3-large"}
|
||||
)
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.rag.ingestion.s3_vectors_ingestion.litellm.aembedding",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_aembedding:
|
||||
mock_aembedding.return_value = SimpleNamespace(
|
||||
data=[{"embedding": [0.1, 0.2, 0.3]}]
|
||||
)
|
||||
|
||||
embeddings = await ingestion.embed(["hello"])
|
||||
|
||||
mock_aembedding.assert_awaited_once_with(
|
||||
model="text-embedding-3-large", input=["hello"]
|
||||
)
|
||||
assert embeddings == [[0.1, 0.2, 0.3]]
|
||||
|
||||
|
||||
def test_s3_vectors_ingestion_prefers_top_level_embedding_config():
|
||||
ingestion = S3VectorsRAGIngestion(
|
||||
ingest_options=_s3_vectors_ingest_options(
|
||||
embedding={"model": "text-embedding-3-small"},
|
||||
vector_store={"embedding_model": "text-embedding-3-large"},
|
||||
)
|
||||
)
|
||||
|
||||
assert ingestion.embedding_config == {"model": "text-embedding-3-small"}
|
||||
Loading…
Add table
Reference in a new issue