diff --git a/tests/test_litellm/proxy/rag_endpoints/test_rag_endpoints.py b/tests/test_litellm/proxy/rag_endpoints/test_rag_endpoints.py index abbf6892a98..d9cb14fe889 100644 --- a/tests/test_litellm/proxy/rag_endpoints/test_rag_endpoints.py +++ b/tests/test_litellm/proxy/rag_endpoints/test_rag_endpoints.py @@ -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, ): diff --git a/ui/litellm-dashboard/src/components/networking.test.ts b/ui/litellm-dashboard/src/components/networking.test.ts index 71296bc1dfd..b254bd23ddd 100644 --- a/ui/litellm-dashboard/src/components/networking.test.ts +++ b/ui/litellm-dashboard/src/components/networking.test.ts @@ -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) => { + const mockFetch = vi.fn().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"); + }); +});