feat: integrate spring ai vectorstore fallback
This commit is contained in:
@@ -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;
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user