feat: integrate spring ai vectorstore fallback

This commit is contained in:
aruo
2026-07-05 10:20:29 +08:00
parent b9ec07de57
commit 5c71f5fc79
12 changed files with 541 additions and 42 deletions
@@ -1,5 +1,8 @@
package com.superbiz.agent.service;
import com.fasterxml.jackson.core.JsonProcessingException;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.superbiz.agent.constant.MilvusConstants;
import io.milvus.client.MilvusServiceClient;
import io.milvus.grpc.SearchResults;
import io.milvus.param.R;
@@ -7,19 +10,26 @@ import io.milvus.param.dml.SearchParam;
import io.milvus.response.SearchResultsWrapper;
import lombok.Getter;
import lombok.Setter;
import com.superbiz.agent.constant.MilvusConstants;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
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.beans.factory.annotation.Autowired;
import org.springframework.beans.factory.annotation.Value;
import org.springframework.stereotype.Service;
import java.util.ArrayList;
import java.util.Collections;
import java.util.List;
import java.util.Map;
/**
* 向量搜索服务
* 负责从 Milvus 中搜索相似向量
* Vector retrieval facade used by lookup_knowledge.
*
* <p>The public API stays stable while the implementation can route to Spring AI
* VectorStore, the original Milvus SDK path, or automatic fallback.</p>
*/
@Service
public class VectorSearchService {
@@ -32,34 +42,84 @@ public class VectorSearchService {
@Autowired
private VectorEmbeddingService embeddingService;
/**
* 搜索相似文档
*
* @param query 查询文本
* @param topK 返回最相似的K个结果
* @return 搜索结果列表
*/
@Autowired
private ObjectProvider<VectorStore> vectorStoreProvider;
@Autowired
private ObjectMapper objectMapper;
@Value("${retrieval.vector-store.mode:auto}")
private String vectorStoreMode = "auto";
@Value("${retrieval.normalization.max-l2-distance:2.0}")
private double maxL2Distance = 2.0;
public List<SearchResult> searchSimilarDocuments(String query, int topK) {
return searchSimilarDocuments(query, topK, null);
}
/**
* 搜索相似文档(支持类别过滤)
*
* @param query 查询文本
* @param topK 返回最相似的K个结果
* @param category 类别过滤(可选,null 表示不过滤)
* @return 搜索结果列表
*/
public List<SearchResult> searchSimilarDocuments(String query, int topK, String category) {
String mode = vectorStoreMode == null ? "auto" : vectorStoreMode.trim().toLowerCase();
return switch (mode) {
case "sdk" -> searchSimilarDocumentsWithSdk(query, topK, category);
case "spring-ai" -> searchSimilarDocumentsWithVectorStore(query, topK, category);
case "auto" -> searchWithAutoFallback(query, topK, category);
default -> {
logger.warn("Unknown retrieval.vector-store.mode={}, using auto mode", vectorStoreMode);
yield searchWithAutoFallback(query, topK, category);
}
};
}
private List<SearchResult> searchWithAutoFallback(String query, int topK, String category) {
try {
logger.info("开始搜索相似文档, 查询: {}, topK: {}, 类别: {}", query, topK, category);
return searchSimilarDocumentsWithVectorStore(query, topK, category);
} catch (Exception e) {
logger.warn("Spring AI VectorStore retrieval failed, falling back to Milvus SDK: {}", e.getMessage());
return searchSimilarDocumentsWithSdk(query, topK, category);
}
}
List<SearchResult> searchSimilarDocumentsWithVectorStore(String query, int topK, String category) {
VectorStore vectorStore = vectorStoreProvider != null ? vectorStoreProvider.getIfAvailable() : null;
if (vectorStore == null) {
throw new IllegalStateException("Spring AI VectorStore bean is unavailable");
}
logger.info("Starting Spring AI VectorStore search: query={}, topK={}, category={}", query, topK, category);
SearchRequest.Builder builder = SearchRequest.builder()
.query(query)
.topK(topK)
.similarityThresholdAll();
if (category != null && !category.trim().isEmpty()) {
String filterExpression = "category == '" + escapeFilterValue(category.trim()) + "'";
builder.filterExpression(filterExpression);
logger.info("Spring AI VectorStore category filter: {}", filterExpression);
}
List<Document> documents = vectorStore.similaritySearch(builder.build());
List<SearchResult> results = new ArrayList<>();
for (Document document : documents) {
SearchResult result = new SearchResult();
result.setId(document.getId());
result.setContent(document.getText());
result.setMetadata(toJson(document.getMetadata()));
result.setRawScore(document.getScore());
result.setScoreLabel("similarity");
result.setScore(toCompatibleL2Distance(document.getScore()));
results.add(result);
}
logger.info("Spring AI VectorStore search complete, candidates={}", results.size());
return results;
}
List<SearchResult> searchSimilarDocumentsWithSdk(String query, int topK, String category) {
try {
logger.info("Starting Milvus SDK search: query={}, topK={}, category={}", query, topK, category);
// 1. 将查询文本向量化
List<Float> queryVector = embeddingService.generateQueryVector(query);
logger.debug("查询向量生成成功, 维度: {}", queryVector.size());
logger.debug("Query vector generated, dimension={}", queryVector.size());
// 2. 构建搜索参数
SearchParam.Builder searchParamBuilder = SearchParam.newBuilder()
.withCollectionName(MilvusConstants.MILVUS_COLLECTION_NAME)
.withVectorFieldName("vector")
@@ -69,33 +129,27 @@ public class VectorSearchService {
.withOutFields(List.of("id", "content", "metadata"))
.withParams("{\"nprobe\":10}");
// 添加类别过滤
if (category != null && !category.trim().isEmpty()) {
String expr = String.format("metadata[\"category\"] == \"%s\"", category);
searchParamBuilder.withExpr(expr);
logger.info("添加类别过滤: {}", expr);
logger.info("Milvus SDK category filter: {}", expr);
}
SearchParam searchParam = searchParamBuilder.build();
// 3. 执行搜索
R<SearchResults> searchResponse = milvusClient.search(searchParam);
R<SearchResults> searchResponse = milvusClient.search(searchParamBuilder.build());
if (searchResponse.getStatus() != 0) {
throw new RuntimeException("向量搜索失败: " + searchResponse.getMessage());
throw new RuntimeException("Vector search failed: " + searchResponse.getMessage());
}
// 4. 解析搜索结果
SearchResultsWrapper wrapper = new SearchResultsWrapper(searchResponse.getData().getResults());
List<SearchResult> results = new ArrayList<>();
for (int i = 0; i < wrapper.getRowRecords(0).size(); i++) {
SearchResult result = new SearchResult();
result.setId((String) wrapper.getIDScore(0).get(i).get("id"));
result.setContent((String) wrapper.getFieldData("content", 0).get(i));
result.setScore(wrapper.getIDScore(0).get(i).getScore());
result.setRawScore((double) result.getScore());
result.setScoreLabel("l2_distance");
// 解析 metadata
Object metadataObj = wrapper.getFieldData("metadata", 0).get(i);
if (metadataObj != null) {
result.setMetadata(metadataObj.toString());
@@ -104,25 +158,50 @@ public class VectorSearchService {
results.add(result);
}
logger.info("搜索完成, 找到 {} 个相似文档", results.size());
logger.info("Milvus SDK search complete, candidates={}", results.size());
return results;
} catch (Exception e) {
logger.error("搜索相似文档失败", e);
throw new RuntimeException("搜索失败: " + e.getMessage(), e);
logger.error("Milvus SDK vector search failed", e);
throw new RuntimeException("Vector search failed: " + e.getMessage(), e);
}
}
/**
* 搜索结果类
*/
private float toCompatibleL2Distance(Double similarity) {
if (similarity == null) {
return (float) maxL2Distance;
}
double bounded = Math.max(0.0, Math.min(1.0, similarity));
return (float) ((1.0 - bounded) * maxL2Distance);
}
private String toJson(Map<String, Object> metadata) {
if (metadata == null || metadata.isEmpty()) {
return null;
}
try {
return objectMapper.writeValueAsString(metadata);
} catch (JsonProcessingException e) {
return metadata.toString();
}
}
private String escapeFilterValue(String value) {
return value.replace("'", "\\'");
}
@Setter
@Getter
public static class SearchResult {
private String id;
private String content;
/**
* Compatibility score used by existing lookup relevance normalization.
* SDK mode keeps L2 distance; VectorStore mode maps similarity into a
* L2-like distance using retrieval.normalization.max-l2-distance.
*/
private float score;
private Double rawScore;
private String scoreLabel;
private String metadata;
}
}
+26
View File
@@ -93,6 +93,30 @@ spring:
min-idle: 0
ai:
vectorstore:
type: milvus
milvus:
initialize-schema: false
database-name: ${milvus.database}
collection-name: business_knowledge
embedding-dimension: ${milvus.vector-dim}
index-type: IVF_FLAT
metric-type: L2
index-parameters: '{"nlist":128}'
id-field-name: id
auto-id: false
content-field-name: content
metadata-field-name: metadata
embedding-field-name: vector
client:
host: ${milvus.host}
port: ${milvus.port}
token: ${milvus.token}
username: ${milvus.username}
password: ${milvus.password}
secure: ${milvus.secure}
connect-timeout-ms: ${milvus.timeout}
# --- Chat: DeepSeek (原生) ---
deepseek:
api-key: sk-1f44696abe644bd684f09cc43f12c557
@@ -133,6 +157,8 @@ rag:
# 检索归一化配置
retrieval:
vector-store:
mode: auto # auto | spring-ai | sdk
normalization:
max-l2-distance: 2.0 # L2 距离上界(BGE-M3 单位向量 = 2.0)
highly-relevant-threshold: 0.75 # similarity >= 0.75 → HIGHLY_RELEVANT
@@ -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;
}
}