feat: integrate spring ai vectorstore fallback
This commit is contained in:
@@ -0,0 +1,148 @@
|
||||
package com.superbiz.agent.service;
|
||||
|
||||
import com.fasterxml.jackson.databind.ObjectMapper;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.mockito.ArgumentCaptor;
|
||||
import org.springframework.ai.document.Document;
|
||||
import org.springframework.ai.vectorstore.SearchRequest;
|
||||
import org.springframework.ai.vectorstore.VectorStore;
|
||||
import org.springframework.beans.factory.ObjectProvider;
|
||||
import org.springframework.test.util.ReflectionTestUtils;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
import static org.junit.jupiter.api.Assertions.assertEquals;
|
||||
import static org.junit.jupiter.api.Assertions.assertTrue;
|
||||
import static org.mockito.ArgumentMatchers.any;
|
||||
import static org.mockito.ArgumentMatchers.eq;
|
||||
import static org.mockito.Mockito.doReturn;
|
||||
import static org.mockito.Mockito.mock;
|
||||
import static org.mockito.Mockito.never;
|
||||
import static org.mockito.Mockito.spy;
|
||||
import static org.mockito.Mockito.verify;
|
||||
import static org.mockito.Mockito.when;
|
||||
|
||||
class VectorSearchServiceTest {
|
||||
|
||||
@Test
|
||||
void sdkModeBypassesVectorStore() {
|
||||
VectorSearchService service = spy(new VectorSearchService());
|
||||
setMode(service, "sdk");
|
||||
VectorSearchService.SearchResult expected = result("sdk-doc", 0.2f);
|
||||
doReturn(List.of(expected))
|
||||
.when(service).searchSimilarDocumentsWithSdk("query", 3, null);
|
||||
|
||||
List<VectorSearchService.SearchResult> results = service.searchSimilarDocuments("query", 3, null);
|
||||
|
||||
assertEquals(List.of(expected), results);
|
||||
verify(service, never()).searchSimilarDocumentsWithVectorStore(any(), eq(3), any());
|
||||
}
|
||||
|
||||
@Test
|
||||
void autoModeUsesVectorStoreWhenAvailable() {
|
||||
VectorStore vectorStore = mock(VectorStore.class);
|
||||
ObjectProvider<VectorStore> provider = mock(ObjectProvider.class);
|
||||
when(provider.getIfAvailable()).thenReturn(vectorStore);
|
||||
when(vectorStore.similaritySearch(any(SearchRequest.class))).thenReturn(List.of(
|
||||
Document.builder()
|
||||
.id("spring-doc")
|
||||
.text("spring content")
|
||||
.metadata(Map.of("_source", "spring.md", "category", "api"))
|
||||
.score(0.8)
|
||||
.build()
|
||||
));
|
||||
|
||||
VectorSearchService service = new VectorSearchService();
|
||||
setMode(service, "auto");
|
||||
setVectorStore(service, provider);
|
||||
|
||||
List<VectorSearchService.SearchResult> results = service.searchSimilarDocuments("query", 3, null);
|
||||
|
||||
assertEquals(1, results.size());
|
||||
assertEquals("spring-doc", results.get(0).getId());
|
||||
assertEquals("similarity", results.get(0).getScoreLabel());
|
||||
assertEquals(0.8, results.get(0).getRawScore(), 0.0001);
|
||||
assertEquals(0.4f, results.get(0).getScore(), 0.0001);
|
||||
assertTrue(results.get(0).getMetadata().contains("spring.md"));
|
||||
}
|
||||
|
||||
@Test
|
||||
void autoModeFallsBackToSdkWhenVectorStoreFails() {
|
||||
VectorStore vectorStore = mock(VectorStore.class);
|
||||
ObjectProvider<VectorStore> provider = mock(ObjectProvider.class);
|
||||
when(provider.getIfAvailable()).thenReturn(vectorStore);
|
||||
when(vectorStore.similaritySearch(any(SearchRequest.class))).thenThrow(new RuntimeException("vectorstore down"));
|
||||
|
||||
VectorSearchService service = spy(new VectorSearchService());
|
||||
setMode(service, "auto");
|
||||
setVectorStore(service, provider);
|
||||
VectorSearchService.SearchResult fallback = result("sdk-doc", 0.3f);
|
||||
doReturn(List.of(fallback))
|
||||
.when(service).searchSimilarDocumentsWithSdk("query", 3, null);
|
||||
|
||||
List<VectorSearchService.SearchResult> results = service.searchSimilarDocuments("query", 3, null);
|
||||
|
||||
assertEquals(List.of(fallback), results);
|
||||
}
|
||||
|
||||
@Test
|
||||
void autoModeFallsBackToSdkWhenVectorStoreUnavailable() {
|
||||
ObjectProvider<VectorStore> provider = mock(ObjectProvider.class);
|
||||
when(provider.getIfAvailable()).thenReturn(null);
|
||||
|
||||
VectorSearchService service = spy(new VectorSearchService());
|
||||
setMode(service, "auto");
|
||||
setVectorStore(service, provider);
|
||||
VectorSearchService.SearchResult fallback = result("sdk-doc", 0.3f);
|
||||
doReturn(List.of(fallback))
|
||||
.when(service).searchSimilarDocumentsWithSdk("query", 3, null);
|
||||
|
||||
List<VectorSearchService.SearchResult> results = service.searchSimilarDocuments("query", 3, null);
|
||||
|
||||
assertEquals(List.of(fallback), results);
|
||||
}
|
||||
|
||||
@Test
|
||||
void vectorStoreSearchUsesCategoryFilter() {
|
||||
VectorStore vectorStore = mock(VectorStore.class);
|
||||
ObjectProvider<VectorStore> provider = mock(ObjectProvider.class);
|
||||
when(provider.getIfAvailable()).thenReturn(vectorStore);
|
||||
when(vectorStore.similaritySearch(any(SearchRequest.class))).thenReturn(List.of());
|
||||
VectorSearchService service = new VectorSearchService();
|
||||
setMode(service, "spring-ai");
|
||||
setVectorStore(service, provider);
|
||||
|
||||
service.searchSimilarDocuments("query", 5, "api");
|
||||
|
||||
ArgumentCaptor<SearchRequest> requestCaptor = ArgumentCaptor.forClass(SearchRequest.class);
|
||||
verify(vectorStore).similaritySearch(requestCaptor.capture());
|
||||
SearchRequest request = requestCaptor.getValue();
|
||||
assertEquals("query", request.getQuery());
|
||||
assertEquals(5, request.getTopK());
|
||||
assertTrue(request.hasFilterExpression());
|
||||
assertTrue(request.toString().contains("category"));
|
||||
assertTrue(request.toString().contains("api"));
|
||||
}
|
||||
|
||||
private static void setMode(VectorSearchService service, String mode) {
|
||||
ReflectionTestUtils.setField(service, "vectorStoreMode", mode);
|
||||
}
|
||||
|
||||
private static void setVectorStore(VectorSearchService service, ObjectProvider<VectorStore> provider) {
|
||||
ReflectionTestUtils.setField(service, "vectorStoreProvider", provider);
|
||||
ReflectionTestUtils.setField(service, "objectMapper", new ObjectMapper());
|
||||
ReflectionTestUtils.setField(service, "maxL2Distance", 2.0);
|
||||
}
|
||||
|
||||
private static VectorSearchService.SearchResult result(String id, float score) {
|
||||
VectorSearchService.SearchResult result = new VectorSearchService.SearchResult();
|
||||
result.setId(id);
|
||||
result.setScore(score);
|
||||
result.setRawScore((double) score);
|
||||
result.setScoreLabel("l2_distance");
|
||||
result.setContent("content");
|
||||
result.setMetadata("{}");
|
||||
return result;
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user