Spring AI RAG 知识库系统搭建
·
一、项目概述
以下是一个基于 Spring AI 框架构建的 RAG(Retrieval-Augmented Generation,检索增强生成) 知识库系统。实现了文档向量化存储、混合检索(向量检索 + 全文检索)、智能问答等核心功能,通过结合大语言模型和知识库检索,提供准确、可溯源的智能问答服务。
1.1 核心特性
- 混合检索架构:同时支持 Milvus 向量检索和 Elasticsearch 全文检索
- RRF 结果融合:使用 Reciprocal Rank Fusion 算法融合多路检索结果
- 流式响应:支持 SSE 流式输出,提升用户体验
- 文档处理:支持 DOCX 文档解析、切片、向量化入库
1.2 技术栈
| 技术组件 | 版本 | 用途 |
|---|---|---|
| Spring Boot | 3.5.5 | 基础框架 |
| Spring AI | 1.0.1 | AI 能力集成 |
| Milvus | 2.6.2 | 向量数据库 |
| Elasticsearch | 7.5.1 | 全文搜索引擎 |
| MySQL | 8.0.33 | 对话记忆存储 |
| LLM | 豆包 doubao-seed-2-0 | 大语言模型 |
| Embedding | BGE-M3 | 文本嵌入模型 |
二、项目架构
2.1 整体架构图
┌─────────────────────────────────────────────────────────────────┐
│ 客户端请求 │
└─────────────────────────────────────────────────────────────────┘
│
▼
┌─────────────────────────────────────────────────────────────────┐
│ Controller 层 │
│ ┌──────────────────┐ ┌──────────────────────────────────────┐ │
│ │ ChatController │ │ RagController │ │
│ │ (普通对话) │ │ /ai/chat, /ai/rag/chat, /ai/rag/* │ │
│ └──────────────────┘ └──────────────────────────────────────┘ │
└─────────────────────────────────────────────────────────────────┘
│
▼
┌─────────────────────────────────────────────────────────────────┐
│ Service 层 │
│ ┌────────────────────────────────────────────────────────────┐ │
│ │ HybridSearchService │ │
│ │ (混合检索 + RRF 融合) │ │
│ └────────────────────────────────────────────────────────────┘ │
│ │ │ │
│ ▼ ▼ │
│ ┌──────────────────────┐ ┌──────────────────────────────┐ │
│ │ MilvusVectorService │ │ ElasticsearchService │ │
│ │ (向量检索) │ │ (全文检索) │ │
│ └──────────────────────┘ └──────────────────────────────┘ │
└─────────────────────────────────────────────────────────────────┘
│
▼
┌─────────────────────────────────────────────────────────────────┐
│ Config 层 │
│ ┌─────────────┐ ┌────────────────┐ ┌────────────────────────┐ │
│ │ ChatConfig │ │ EmbeddingConfig│ │ MilvusProperties │ │
│ │ (LLM配置) │ │ (嵌入模型配置) │ │ ElasticsearchProperties│ │
│ └─────────────┘ └────────────────┘ └────────────────────────┘ │
└─────────────────────────────────────────────────────────────────┘
│
▼
┌─────────────────────────────────────────────────────────────────┐
│ 外部服务/存储 │
│ ┌───────────┐ ┌─────────────┐ ┌───────────┐ ┌──────────────┐ │
│ │ Milvus │ │Elasticsearch│ │ MySQL │ │ 豆包/BGE-M3 │ │
│ │ (向量库) │ │ (全文索引) │ │(对话记忆) │ │ (AI模型) │ │
│ └───────────┘ └─────────────┘ └───────────┘ └──────────────┘ │
└─────────────────────────────────────────────────────────────────┘
2.2 目录结构
spring-ai/
├── src/main/java/com/ai/springai/
│ ├── SpringAiApplication.java # 应用启动类
│ ├── config/ # 配置层
│ │ ├── ChatConfig.java # ChatClient 配置
│ │ ├── EmbeddingConfig.java # 嵌入模型配置
│ │ ├── MilvusProperties.java # Milvus 属性配置
│ │ └── ElasticsearchProperties.java # ES 属性配置
│ ├── controller/ # 控制器层
│ │ ├── ChatController.java # 普通对话控制器
│ │ └── RagController.java # RAG 对话控制器
│ ├── model/ # 数据模型
│ │ ├── Document.java # 文档实体
│ │ └── SearchResult.java # 检索结果实体
│ └── service/ # 服务层
│ ├── HybridSearchService.java # 混合检索服务
│ ├── MilvusVectorService.java # Milvus 向量服务
│ └── ElasticsearchService.java # ES 全文检索服务
├── src/main/resources/
│ └── application.yml # 应用配置文件
三、核心模块详解
3.1 配置模块
3.1.1 application.yml 完整配置
server:
port: 8888
spring:
ai:
openai:
base-url: https://ark.cn-beijing.volces.com/api/v3
api-key: ${API_KEY}
chat:
options:
model: doubao-seed-2-0-code-preview-260215
embedding:
base-url: http://localhost:7997
options:
model: bge-m3
chat:
memory:
repository:
jdbc:
schema: classpath:schema-@@platform@@.sql
initialize-schema: never
datasource:
driver-class-name: com.mysql.cj.jdbc.Driver
url: jdbc:mysql://localhost:3306/knowledge_base
username: root
password: ${DB_PASSWORD}
milvus:
host: http://127.0.0.1
port: 19530
database: default
collection: knowledge_base
embedding-dimension: 1024
elasticsearch:
host: 127.0.0.1
port: 9200
index: knowledge_base
3.1.2 ChatConfig - 大模型配置
@Configuration
public class ChatConfig {
@Bean
public OpenAiApi openAiApi() {
return OpenAiApi.builder()
.baseUrl(baseUrl) // 火山引擎 API 地址
.apiKey(apiKey)
.completionsPath("/chat/completions")
.build();
}
@Bean
public ChatClient chatClient(ChatModel chatModel) {
return ChatClient.builder(chatModel)
.defaultSystem("你是一个智能助手,帮助用户解决问题")
.build();
}
}
关键点:
- 使用火山引擎豆包模型,兼容 OpenAI API 协议
- 配置默认系统提示词,设定 AI 助手角色
3.1.3 EmbeddingConfig - 嵌入模型配置
@Configuration
public class EmbeddingConfig {
@Value("${spring.ai.openai.embedding.base-url}")
private String embeddingBaseUrl;
@Value("${spring.ai.openai.embedding.options.model:bge-m3}")
private String model;
@Bean
public EmbeddingModel embeddingModel() {
return new CustomEmbeddingModel(embeddingBaseUrl, model);
}
private static class CustomEmbeddingModel implements EmbeddingModel {
private final String baseUrl;
private final String model;
private final ObjectMapper objectMapper;
public CustomEmbeddingModel(String baseUrl, String model) {
this.baseUrl = baseUrl;
this.model = model;
this.objectMapper = new ObjectMapper();
}
@SneakyThrows
@Override
public EmbeddingResponse call(EmbeddingRequest request) {
List<String> inputs = request.getInstructions();
Map<String, Object> requestBody = Map.of("input", inputs);
String jsonBody = objectMapper.writeValueAsString(requestBody);
String responseBody = Unirest.post(baseUrl + "/embeddings")
.header("Content-Type", "application/json")
.header("Accept", "*/*")
.body(jsonBody)
.asString()
.getBody();
Map<String, Object> response = objectMapper.readValue(responseBody, Map.class);
List<Map<String, Object>> dataList =
(List<Map<String, Object>>) response.get("data");
List<Embedding> embeddings = new ArrayList<>();
for (int i = 0; i < dataList.size(); i++) {
List<Number> embeddingData =
(List<Number>) dataList.get(i).get("embedding");
float[] embeddingArray = new float[embeddingData.size()];
for (int j = 0; j < embeddingData.size(); j++) {
embeddingArray[j] = embeddingData.get(j).floatValue();
}
embeddings.add(new Embedding(embeddingArray, i));
}
return new EmbeddingResponse(embeddings);
}
@Override
public float[] embed(String text) {
EmbeddingResponse response = call(new EmbeddingRequest(List.of(text), null));
return response.getResults().get(0).getOutput();
}
@Override
public float[] embed(Document document) {
return embed(document.getText());
}
@Override
public List<float[]> embed(List<String> texts) {
EmbeddingResponse response = call(new EmbeddingRequest(texts, null));
return response.getResults().stream()
.map(Embedding::getOutput)
.collect(java.util.stream.Collectors.toList());
}
}
}
关键点:
- 自定义
CustomEmbeddingModel实现 Spring AI 的EmbeddingModel接口 - 调用本地 BGE-M3 嵌入服务(localhost:7997)
- 支持批量文本嵌入
3.1.4 Milvus 配置类
@ConfigurationProperties(prefix = "milvus")
public class MilvusProperties {
private String host = "127.0.0.1";
private int port = 19530;
private String database = "default";
private String collection = "knowledge_base";
private int embeddingDimension = 1024;
}
3.2 服务模块
3.2.1 HybridSearchService - 混合检索核心
混合搜索 MIlvus 与 ES:
public SearchResult search(String query, int topK, Map<String, Object> filters, boolean enableRerank) {
long startTime = System.currentTimeMillis();
log.info("Starting hybrid search for query: {}, topK: {}, enableRerank: {}", query, topK, enableRerank);
SearchResult milvusResult = milvusVectorService.search(query, topK, filters, false);
log.debug("Milvus search returned {} results", milvusResult.getTotalHits());
SearchResult esResult = elasticsearchService.search(query, topK, filters);
log.debug("Elasticsearch search returned {} results", esResult.getTotalHits());
List<Document> fusedDocuments = rrfFusion(milvusResult.getDocuments(), esResult.getDocuments(), topK);
log.debug("RRF fusion returned {} unique documents", fusedDocuments.size());
long endTime = System.currentTimeMillis();
return SearchResult.builder()
.documents(fusedDocuments)
.query(query)
.totalHits(fusedDocuments.size())
.searchTime((endTime - startTime) / 1000.0)
.build();
}
RRF 融合算法:
private List<Document> rrfFusion(List<Document> vectorDocs, List<Document> textDocs, int topK) {
Map<String, Double> rrfScores = new HashMap<>();
int RrfK = 60;
for (int rank = 0; rank < vectorDocs.size(); rank++) {
double rrfScore = 1.0 / (RRF_K + rank + 1);
rrfScores.merge(docId, rrfScore, Double::sum);
}
for (int rank = 0; rank < textDocs.size(); rank++) {
double rrfScore = 1.0 / (RRF_K + rank + 1);
rrfScores.merge(docId, rrfScore, Double::sum);
}
return sortedByScoreDesc(rrfScores).limit(topK);
}
RRF 公式:
RRF_score(d) = Σ 1/(k + rank(d))
其中 k=60 是平滑参数,rank(d) 是文档在检索结果中的排名。
上下文构建:
public String buildContext(List<Document> documents, int maxContextLength) {
StringBuilder context = new StringBuilder();
context.append("以下是相关的知识库内容:\n\n");
for (Document doc : documents) {
String docContent = String.format(
"【文档%d】(来源: %s, 相关度: %.2f)\n%s\n\n",
i + 1, doc.getSource(), doc.getScore(), doc.getContent()
);
if (currentLength + docContent.length() > maxContextLength) break;
context.append(docContent);
}
return context.toString();
}
3.2.2 MilvusVectorService - 向量检索服务
初始化连接:
@PostConstruct
public void init() {
try {
ConnectConfig connectConfig = ConnectConfig.builder()
.uri(milvusProperties.getHost() + ":" + milvusProperties.getPort())
.dbName(
milvusProperties.getDatabase() == null
? "default"
: milvusProperties.getDatabase()
)
.build();
this.milvusClient = new MilvusClientV2(connectConfig);
log.info("Milvus client initialized successfully at {}:{}",
milvusProperties.getHost(), milvusProperties.getPort());
} catch (Exception e) {
log.error("Failed to initialize Milvus client at {}:{}. Error: {}",
milvusProperties.getHost(), milvusProperties.getPort(), e.getMessage());
}
}
用户语义向量化:
public float[] generateEmbedding(String text) {
log.debug("Generating embedding for text: {}", text.substring(0, Math.min(50, text.length())));
return embeddingModel.embed(text);
}
向量检索:
public SearchResult search(String query, int topK, Map<String, Object> filter, boolean enableRerank) {
long startTime = System.currentTimeMillis();
log.debug("Milvus vector search for query: {}, topK: {}, enableRerank: {}", query, topK, enableRerank);
if (milvusClient == null) {
log.error("Milvus client is not initialized. Check Milvus connection at {}:{}",
milvusProperties.getHost(), milvusProperties.getPort());
return SearchResult.builder()
.documents(new ArrayList<>())
.query(query)
.totalHits(0)
.searchTime(0)
.build();
}
try {
float[] queryEmbedding = generateEmbedding(query);
int retrieveCount = enableRerank ? Math.max(topK * 3, 10) : topK;
SearchReq.SearchReqBuilder searchReqBuilder = SearchReq.builder()
.collectionName(milvusProperties.getCollection())
.data(List.of(new FloatVec(queryEmbedding)))
.topK(retrieveCount)
.outputFields(List.of("content", "metadata", "source"));
if (filter != null && !filter.isEmpty()) {
searchReqBuilder.filter(buildFilterExpression(filter));
}
SearchResp searchResp = milvusClient.search(searchReqBuilder.build());
List<Document> documents = new ArrayList<>();
var searchResults = searchResp.getSearchResults();
for (var resultList : searchResults) {
for (var result : resultList) {
Document doc = Document.builder()
.id(String.valueOf(result.getId()))
.score(result.getScore())
.content((String) result.getEntity().get("content"))
.source((String) result.getEntity().get("source"))
.metadata((JsonObject) result.getEntity().get("metadata"))
.build();
documents.add(doc);
}
}
documents = documents.stream()
.limit(topK)
.collect(Collectors.toList());
long endTime = System.currentTimeMillis();
return SearchResult.builder()
.documents(documents)
.query(query)
.totalHits(documents.size())
.searchTime((endTime - startTime) / 1000.0)
.build();
} catch (Exception e) {
log.error("Milvus search failed", e);
return SearchResult.builder()
.documents(new ArrayList<>())
.query(query)
.totalHits(0)
.searchTime(0)
.build();
}
}
3.2.3 ElasticsearchService - 全文检索服务
索引映射配置:
{
"properties": {
"id": { "type": "keyword" },
"content": {
"type": "text",
"analyzer": "ik_max_word",
"search_analyzer": "ik_smart"
},
"source": { "type": "keyword" },
"metadata": { "type": "object", "enabled": true }
}
}
全文检索:
public SearchResult search(String query, int topK) {
MultiMatchQueryBuilder multiMatchQuery = QueryBuilders
.multiMatchQuery(query)
.field("content", 2.0f);
SearchSourceBuilder sourceBuilder = new SearchSourceBuilder()
.query(multiMatchQuery)
.size(topK);
return esClient.search(searchRequest);
}
3.3 控制器模块
3.3.1 RagController - RAG 对话接口
核心接口:
| 接口 | 方法 | 功能 |
|---|---|---|
/ai/rag/chat | GET | RAG 增强对话(流式) |
RAG 对话流程:
@GetMapping("/ai/rag/chat")
public Flux<String> ragChat(String prompt, String chatId, int topK, boolean enableRerank) {
SearchResult searchResult = hybridSearchService.search(prompt, topK, null, enableRerank);
String context = hybridSearchService.buildContext(searchResult.getDocuments(), MAX_CONTEXT_LENGTH);
String enhancedPrompt = buildEnhancedPrompt(prompt, context);
return chatClient.prompt()
.user(enhancedPrompt)
.stream()
.content();
}
增强提示词模板:
private String buildEnhancedPrompt(String userQuery, String context) {
return """
你是一个智能问答助手。请基于以下知识库内容回答用户的问题。
如果知识库中没有相关信息,请明确告知用户,不要编造答案。
%s
用户问题:%s
请提供准确、详细的回答:
""".formatted(context, userQuery);
}
3.4 数据模型
Document 实体
@Data
@Builder
public class Document {
private String id; // 文档唯一标识
private String content; // 文档内容
private double score; // 相关性得分
private JsonObject metadata; // 元数据
private String source; // 来源文件
}
SearchResult 实体
@Data
@Builder
public class SearchResult {
private List<Document> documents; // 检索结果列表
private String query; // 查询语句
private long totalHits; // 命中总数
private double searchTime; // 检索耗时(秒)
}
四、文档处理流程
4.1 文档导入流程
┌─────────────┐ ┌─────────────┐ ┌─────────────┐ ┌─────────────┐
│ DOCX 文件 │ -> │ 文本提取 │ -> │ 文档切片 │ -> │ 向量化 │
└─────────────┘ └─────────────┘ └─────────────┘ └─────────────┘
│
▼
┌─────────────────────────────────────────────────┐
│ 双写存储 │
│ ┌───────────────┐ ┌───────────────────┐ │
│ │ Milvus │ │ Elasticsearch │ │
│ │ (向量索引) │ │ (全文索引) │ │
│ └───────────────┘ └───────────────────┘ │
└─────────────────────────────────────────────────┘
4.2 切片策略
private static final int CHUNK_SIZE = 500; // 切片大小
private static final int OVERLAP = 50; // 重叠大小
private List<String> splitText(String text, int chunkSize, int overlap) {
text = text.replaceAll("[\\t\\r\\n]+", ""); // 清理空白字符
while (start < length) {
String chunk = text.substring(start, Math.min(start + chunkSize, length));
chunks.add(chunk);
start += chunkSize;
}
return chunks;
}
更多推荐




所有评论(0)