package com.wok.supportbot.rag; import com.baomidou.mybatisplus.core.conditions.query.LambdaQueryWrapper; import com.wok.supportbot.dao.AiModelConfigMapper; import com.wok.supportbot.entity.AiModelConfig; import lombok.extern.slf4j.Slf4j; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.http.*; import org.springframework.stereotype.Service; import org.springframework.web.client.RestTemplate; import java.time.Duration; import java.util.*; import java.util.stream.Collectors; /** * 重排序服务 * 查询 ai_model_config 表中 RERANK 类型的活跃配置,调用对应的 Rerank API 对候选文档精排。 * 支持 DashScope 和 OpenAI 兼容两种协议,无配置或调用异常时 fallback 到原始排序。 */ @Service @Slf4j public class RerankerService { @Autowired private AiModelConfigMapper aiModelConfigMapper; /** HTTP 超时时间(秒) */ private static final int TIMEOUT_SECONDS = 3; /** * 对候选文档执行重排序 * - 有 RERANK 配置 → 调用对应提供商的 Rerank API * - 无配置 / 超时 / 异常 → fallback 返回 candidates 前 topN 条 * * @param query 用户查询 * @param candidates RRF 融合后的候选文档列表 * @param topN 最终返回的文档数量 * @return 按相关性降序排列的 topN 条文档 */ public List rerank(String query, List candidates, int topN) { if (candidates == null || candidates.isEmpty()) { return Collections.emptyList(); } // 查询 RERANK 类型的活跃配置 AiModelConfig config = getActiveRerankConfig(); if (config == null) { log.debug("未找到 RERANK 活跃配置,使用 RRF 原始排序作为 fallback"); return fallback(candidates, topN); } try { List reranked = doRerank(query, candidates, topN, config); log.debug("Reranker 精排完成: provider={}, 输入 {} 条, 输出 {} 条", config.getProvider(), candidates.size(), reranked.size()); return reranked; } catch (Exception e) { log.warn("Reranker 调用失败 (provider={}): {}, 使用 RRF 原始排序作为 fallback", config.getProvider(), e.getMessage()); return fallback(candidates, topN); } } /** * 获取 RERANK 类型的活跃配置 */ private AiModelConfig getActiveRerankConfig() { LambdaQueryWrapper wrapper = new LambdaQueryWrapper<>(); wrapper.eq(AiModelConfig::getAppType, "RERANK") .eq(AiModelConfig::getIsActive, true) .last("LIMIT 1"); return aiModelConfigMapper.selectOne(wrapper); } /** * 根据提供商类型分发 Rerank API 调用 */ private List doRerank(String query, List candidates, int topN, AiModelConfig config) { String provider = config.getProvider().toLowerCase(); if ("dashscope".equals(provider)) { return dashscopeRerank(query, candidates, topN, config); } else { return openaiCompatibleRerank(query, candidates, topN, config); } } /** * DashScope Rerank API 调用 * POST https://dashscope.aliyuncs.com/api/v1/services/rerank */ private List dashscopeRerank(String query, List candidates, int topN, AiModelConfig config) { String url = "https://dashscope.aliyuncs.com/api/v1/services/rerank"; List documents = candidates.stream() .map(RrfFusion.ScoredDocument::getContent) .collect(Collectors.toList()); // 构建请求体 Map parameters = new LinkedHashMap<>(); parameters.put("top_n", topN); parameters.put("return_documents", true); Map input = new LinkedHashMap<>(); input.put("query", query); input.put("documents", documents); Map body = new LinkedHashMap<>(); body.put("model", config.getModelName() != null ? config.getModelName() : "gte-rerank"); body.put("input", input); body.put("parameters", parameters); // 发送请求 HttpHeaders headers = new HttpHeaders(); headers.setContentType(MediaType.APPLICATION_JSON); headers.setBearerAuth(config.getApiKey()); ResponseEntity response = postWithTimeout(url, headers, body); // 解析响应:{"output":{"results":[{"index":0,"relevance_score":0.95},...]}} Map responseBody = response.getBody(); if (responseBody == null) { throw new RuntimeException("DashScope Rerank 响应为空"); } Map output = (Map) responseBody.get("output"); if (output == null) { throw new RuntimeException("DashScope Rerank 响应缺少 output 字段"); } List results = (List) output.get("results"); if (results == null || results.isEmpty()) { throw new RuntimeException("DashScope Rerank 结果为空"); } return mapResultsToDocs(results, candidates); } /** * OpenAI 兼容 Rerank API 调用 * POST {baseUrl}/rerank */ private List openaiCompatibleRerank(String query, List candidates, int topN, AiModelConfig config) { String baseUrl = config.getBaseUrl(); if (baseUrl == null || baseUrl.isBlank()) { throw new IllegalArgumentException( "OpenAI 兼容 Rerank 提供商 [" + config.getProvider() + "] 未配置 baseUrl"); } // 确保 URL 拼接正确 String url = baseUrl.endsWith("/") ? baseUrl + "rerank" : baseUrl + "/rerank"; List documents = candidates.stream() .map(RrfFusion.ScoredDocument::getContent) .collect(Collectors.toList()); // 构建请求体 Map body = new LinkedHashMap<>(); body.put("model", config.getModelName()); body.put("query", query); body.put("documents", documents); body.put("top_n", topN); // 发送请求 HttpHeaders headers = new HttpHeaders(); headers.setContentType(MediaType.APPLICATION_JSON); headers.setBearerAuth(config.getApiKey()); ResponseEntity response = postWithTimeout(url, headers, body); // 解析响应:{"results":[{"index":0,"relevance_score":0.95},...]} Map responseBody = response.getBody(); if (responseBody == null) { throw new RuntimeException("OpenAI 兼容 Rerank 响应为空"); } List results = (List) responseBody.get("results"); if (results == null || results.isEmpty()) { throw new RuntimeException("OpenAI 兼容 Rerank 结果为空"); } return mapResultsToDocs(results, candidates); } /** * 将 Rerank API 返回的结果映射为 ScoredDocument 列表 * * @param results API 返回的 results 数组,每项含 index 和 relevance_score * @param candidates 原始候选文档列表(用于通过 index 关联文档内容) * @return 按 relevance_score 降序排列的文档列表 */ private List mapResultsToDocs(List results, List candidates) { List reranked = new ArrayList<>(results.size()); for (Map result : results) { int index = ((Number) result.get("index")).intValue(); double score = ((Number) result.get("relevance_score")).doubleValue(); if (index >= 0 && index < candidates.size()) { RrfFusion.ScoredDocument original = candidates.get(index); reranked.add(new RrfFusion.ScoredDocument( original.getId(), original.getContent(), original.getMetadata(), score )); } } // 按 relevance_score 降序排列 reranked.sort((a, b) -> Double.compare(b.getScore(), a.getScore())); return reranked; } /** * 带超时的 HTTP POST 请求 */ @SuppressWarnings("unchecked") private ResponseEntity postWithTimeout(String url, HttpHeaders headers, Map body) { RestTemplate restTemplate = new RestTemplate(); HttpEntity> entity = new HttpEntity<>(body, headers); // RestTemplate 默认无超时,此处依赖连接/读取超时由底层控制 // Spring Boot 3.x 中可使用 RestClient 替代以获得更好的超时支持 return restTemplate.exchange(url, HttpMethod.POST, entity, Map.class); } /** * Fallback:直接截取前 topN 条候选文档 */ private List fallback(List candidates, int topN) { int limit = Math.min(topN, candidates.size()); return new ArrayList<>(candidates.subList(0, limit)); } }