mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
test(vector-stores): cover S3 ingest embedding model
This commit is contained in:
parent
4ba8517134
commit
44f2a63c09
2 changed files with 88 additions and 0 deletions
|
|
@ -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,
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
});
|
||||
});
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue