test(vector-stores): cover S3 ingest embedding model

This commit is contained in:
Yujong Lee 2026-08-30 17:52:09 -07:00
parent 4ba8517134
commit 44f2a63c09
2 changed files with 88 additions and 0 deletions

View file

@ -6,6 +6,7 @@ Covers:
"""
import io
import json
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -47,6 +48,41 @@ def client_internal_user():
app.dependency_overrides = original_overrides
@pytest.mark.asyncio
async def test_s3_rag_ingest_persists_embedding_model_for_managed_search():
from litellm.proxy.rag_endpoints.endpoints import (
_save_vector_store_to_db_from_rag_ingest,
)
table = MagicMock()
table.find_unique = AsyncMock(return_value=None)
created_row = MagicMock()
created_row.model_dump.return_value = {
"vector_store_id": "documents:index",
"custom_llm_provider": "s3_vectors",
}
table.create = AsyncMock(return_value=created_row)
prisma_client = MagicMock()
prisma_client.db.litellm_managedvectorstorestable = table
await _save_vector_store_to_db_from_rag_ingest(
response={"vector_store_id": "documents:index"},
ingest_options={
"embedding": {"model": "text-embedding-3-large"},
"vector_store": {
"custom_llm_provider": "s3_vectors",
"vector_bucket_name": "documents",
},
},
prisma_client=prisma_client,
user_api_key_dict=UserAPIKeyAuth(user_id="user-1", team_id="team-1"),
)
persisted_params = json.loads(table.create.await_args.kwargs["data"]["litellm_params"])
assert persisted_params["embedding_model"] == "text-embedding-3-large"
assert persisted_params["vector_bucket_name"] == "documents"
def test_internal_user_viewer_rag_ingest_without_vector_store_id_rejected(
client_internal_user_viewer,
):

View file

@ -806,3 +806,55 @@ describe("daily activity api_key filter", () => {
expect(requestedUrl(mockFetch)).toContain("user_id=");
});
});
describe("ragIngestCall", () => {
const originalFetch = global.fetch;
afterEach(() => {
global.fetch = originalFetch;
});
const sendRequest = async (provider: string, providerParams: Record<string, unknown>) => {
const mockFetch = vi.fn<typeof fetch>().mockResolvedValue(
new Response(JSON.stringify({ status: "completed" }), {
status: 200,
headers: { "Content-Type": "application/json" },
}),
);
global.fetch = mockFetch;
await Networking.ragIngestCall(
"sk-key",
new File(["content"], "document.txt"),
provider,
undefined,
undefined,
undefined,
providerParams,
);
const body = mockFetch.mock.calls[0][1]?.body as FormData;
return JSON.parse(String(body.get("request")));
};
it("serializes embedding_model at the provider-defined ingest level", async () => {
const s3Request = await sendRequest("s3_vectors", {
vector_bucket_name: "documents",
embedding_model: "text-embedding-3-large",
});
const bedrockRequest = await sendRequest("bedrock", {
embedding_model: "amazon.titan-embed-text-v2:0",
});
expect(s3Request).toEqual({
ingest_options: {
embedding: { model: "text-embedding-3-large" },
vector_store: {
custom_llm_provider: "s3_vectors",
vector_bucket_name: "documents",
},
},
});
expect(bedrockRequest.ingest_options.vector_store.embedding_model).toBe("amazon.titan-embed-text-v2:0");
});
});