一、项目概述

以下是一个基于 Spring AI 框架构建的 RAG(Retrieval-Augmented Generation,检索增强生成) 知识库系统。实现了文档向量化存储、混合检索(向量检索 + 全文检索)、智能问答等核心功能,通过结合大语言模型和知识库检索,提供准确、可溯源的智能问答服务。

1.1 核心特性

  • 混合检索架构:同时支持 Milvus 向量检索和 Elasticsearch 全文检索
  • RRF 结果融合:使用 Reciprocal Rank Fusion 算法融合多路检索结果
  • 流式响应:支持 SSE 流式输出,提升用户体验
  • 文档处理:支持 DOCX 文档解析、切片、向量化入库

1.2 技术栈

技术组件版本用途
Spring Boot3.5.5基础框架
Spring AI1.0.1AI 能力集成
Milvus2.6.2向量数据库
Elasticsearch7.5.1全文搜索引擎
MySQL8.0.33对话记忆存储
LLM豆包 doubao-seed-2-0大语言模型
EmbeddingBGE-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/chatGETRAG 增强对话(流式)

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;
}
Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐