本地 RAG 知识库
You can not select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
 
 
 
 
 
 

277 lines
11 KiB

package com.wok.supportbot.controller;
import com.wok.supportbot.app.AssistantApp;
import com.wok.supportbot.app.ChatContext;
import com.wok.supportbot.app.SuggestionGenerator;
import com.wok.supportbot.cache.SuggestionCache;
import com.wok.supportbot.config.RoleAccessConfig;
import com.wok.supportbot.entity.ApiKey;
import com.wok.supportbot.entity.SearchResult;
import com.wok.supportbot.rag.CategoryFilter;
import com.wok.supportbot.rag.HybridSearchService;
import com.wok.supportbot.rag.SearchMode;
import com.wok.supportbot.service.ApiKeyService;
import com.wok.supportbot.service.CustomerServiceRoleService;
import com.wok.supportbot.service.CustomerServiceRoleService.RoleScope;
import com.fasterxml.jackson.core.type.TypeReference;
import com.fasterxml.jackson.databind.ObjectMapper;
import jakarta.servlet.http.HttpServletRequest;
import lombok.extern.slf4j.Slf4j;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.http.MediaType;
import org.springframework.http.ResponseEntity;
import org.springframework.util.StringUtils;
import org.springframework.web.bind.annotation.*;
import reactor.core.publisher.Flux;
import java.util.Collections;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
/**
* 开放 API 接口
* 提供第三方系统调用的对话和检索能力,鉴权由 ApiKeyAuthFilter 处理。
*/
@Slf4j
@RestController
@RequestMapping("/open-api")
public class OpenApiController {
@Autowired
private AssistantApp assistantApp;
@Autowired
private HybridSearchService hybridSearchService;
@Autowired
private SuggestionCache suggestionCache;
@Autowired
private SuggestionGenerator suggestionGenerator;
@Autowired
private ApiKeyService apiKeyService;
@Autowired
private CustomerServiceRoleService customerServiceRoleService;
@Autowired
private CategoryFilter categoryFilter;
@Autowired
private RoleAccessConfig roleAccessConfig;
private static final ObjectMapper objectMapper = new ObjectMapper();
/**
* 同步对话接口(已接入 ChatPipeline,补齐角色/RAG/FAQ/MCP/分类隔离能力)
*/
@PostMapping("/chat")
public ResponseEntity<Map<String, Object>> chat(
@RequestParam String message,
@RequestParam(required = false) String roleId,
@RequestParam(required = false) String chatId,
@RequestParam(required = false) String categoryIds,
@RequestParam(required = false) String rewriteStrategy,
@RequestParam(required = false) Boolean enableRag,
HttpServletRequest request) {
try {
ApiKey apiKey = getApiKeyFromRequest(request);
String resolvedChatId = (chatId != null && !chatId.isBlank())
? chatId
: "openapi-" + apiKey.getId() + "-" + System.currentTimeMillis();
ChatContext ctx = buildOpenApiChatContext(message, resolvedChatId, apiKey, roleId,
categoryIds, rewriteStrategy, enableRag, false);
String reply = assistantApp.chat(ctx);
Map<String, Object> result = new HashMap<>();
result.put("success", true);
result.put("data", Map.of(
"reply", reply,
"chatId", resolvedChatId
));
return ResponseEntity.ok(result);
} catch (Exception e) {
log.error("开放 API 对话失败", e);
return ResponseEntity.status(500).body(Map.of(
"success", false,
"message", "对话失败:" + e.getMessage()
));
}
}
/**
* SSE 流式对话接口(已接入 ChatPipeline,补齐角色/RAG/FAQ/MCP/分类隔离能力)
*/
@GetMapping(value = "/chat/stream", produces = MediaType.TEXT_EVENT_STREAM_VALUE)
public Flux<String> chatStream(
@RequestParam String message,
@RequestParam(required = false) String roleId,
@RequestParam(required = false) String chatId,
@RequestParam(required = false) String categoryIds,
@RequestParam(required = false) String rewriteStrategy,
@RequestParam(required = false) Boolean enableRag,
HttpServletRequest request) {
ApiKey apiKey = getApiKeyFromRequest(request);
String resolvedChatId = (chatId != null && !chatId.isBlank())
? chatId
: "openapi-stream-" + apiKey.getId() + "-" + System.currentTimeMillis();
ChatContext ctx = buildOpenApiChatContext(message, resolvedChatId, apiKey, roleId,
categoryIds, rewriteStrategy, enableRag, true);
return assistantApp.chatStreamOpenAi(ctx);
}
/**
* 知识库检索接口
*
* @param query 查询文本
* @param topK 返回条数(默认 5)
* @param searchMode 检索模式:VECTOR / KEYWORD / HYBRID(默认 VECTOR)
* @return 检索结果列表
*/
@PostMapping("/rag/search")
public ResponseEntity<Map<String, Object>> ragSearch(
@RequestParam String query,
@RequestParam(defaultValue = "5") int topK,
@RequestParam(defaultValue = "VECTOR") String searchMode) {
try {
SearchMode mode;
try {
mode = SearchMode.valueOf(searchMode.toUpperCase());
} catch (IllegalArgumentException e) {
mode = SearchMode.VECTOR;
}
List<SearchResult> results = hybridSearchService.search(
query, mode, topK, 0.0, Collections.emptyList());
Map<String, Object> result = new HashMap<>();
result.put("success", true);
result.put("data", results);
result.put("total", results.size());
return ResponseEntity.ok(result);
} catch (Exception e) {
log.error("开放 API 检索失败", e);
return ResponseEntity.status(500).body(Map.of(
"success", false,
"message", "检索失败:" + e.getMessage()
));
}
}
/**
* 获取 AI 推荐问题列表(Open API 用,ApiKeyAuthFilter 守卫)。
*
* @param chatId 会话 ID
* @return { success: true, data: { suggestions: [...] } }
*/
@GetMapping("/suggestions")
public ResponseEntity<Map<String, Object>> getSuggestions(@RequestParam String chatId) {
try {
List<String> suggestions = suggestionGenerator.generate(chatId);
return ResponseEntity.ok(Map.of(
"success", true,
"data", Map.of("suggestions", suggestions)
));
} catch (Exception e) {
log.error("开放 API 获取建议问题失败: chatId={}", chatId, e);
return ResponseEntity.ok(Map.of(
"success", true,
"data", Map.of("suggestions", List.of())
));
}
}
/**
* 从 request attribute 获取已鉴权的 API Key 信息。
*/
private ApiKey getApiKeyFromRequest(HttpServletRequest request) {
Object attr = request.getAttribute("apiKey");
if (attr instanceof ApiKey apiKey) {
return apiKey;
}
throw new IllegalStateException("API Key 鉴权信息缺失");
}
/**
* 构造 Open API 对话的 ChatContext(含角色解析、严格隔离判断、分类过滤)。
*/
private ChatContext buildOpenApiChatContext(String message, String chatId, ApiKey apiKey,
String roleIdStr, String categoryIds, String rewriteStrategy, Boolean enableRag, boolean streaming) {
// 1. 解析 API Key 允许的角色列表
List<Long> allowedRoleIds = parseAllowedRoleIds(apiKey.getRoleIds());
Long roleId = resolveRoleId(roleIdStr, allowedRoleIds);
// 2. 获取角色范围
RoleScope scope = customerServiceRoleService.getRoleScope(roleId);
// 3. 解析系统提示词(角色人设)
String systemPrompt = scope.hasRole() && StringUtils.hasText(scope.systemPrompt())
? scope.systemPrompt() : null;
// 4. 解析分类隔离范围(角色有绑定则强制使用,否则使用客户端传入)
List<Long> catIds = resolveCategoryIds(scope, categoryIds);
// 5. 严格隔离:角色无分类则拒绝检索 KB
boolean useRag = (enableRag == null || enableRag) // 默认启用 RAG
&& !(roleAccessConfig.isStrictIsolation() && scope.hasRole() && scope.categoryIds().isEmpty());
// 6. 未指定策略时默认 MULTI_QUERY
String strategy = (rewriteStrategy != null && !rewriteStrategy.isBlank())
? rewriteStrategy : "MULTI_QUERY";
return new ChatContext(message, chatId, "CHAT", systemPrompt,
scope.hasRole() ? scope.allowedMcpTools() : null,
catIds, strategy, useRag, streaming, roleId, scope.name(), null, apiKey.getId(), null);
}
/**
* 解析 API Key 绑定的角色 ID 列表(JSON 字符串如 "[1,2,3]" → List<Long>)。
* null / 空数组表示不限制(返回空列表)。
*/
private List<Long> parseAllowedRoleIds(String roleIdsJson) {
if (roleIdsJson == null || roleIdsJson.isBlank() || "[]".equals(roleIdsJson.trim())) {
return Collections.emptyList();
}
try {
return objectMapper.readValue(roleIdsJson, new TypeReference<List<Long>>() {});
} catch (Exception e) {
log.warn("解析 API Key role_ids 失败: {}", roleIdsJson, e);
return Collections.emptyList();
}
}
/**
* 从客户端传入的 roleId 参数解析有效角色 ID。
* 若 API Key 绑定了角色列表,则仅允许使用绑定的角色;未绑定时允许任意角色。
* 传入 roleId 不在允许列表内时降级为不指定角色(无角色人设/分类/MCP 权限)。
*/
private Long resolveRoleId(String roleIdStr, List<Long> allowedRoleIds) {
if (roleIdStr == null || roleIdStr.isBlank()) {
return null;
}
Long roleId;
try {
roleId = Long.valueOf(roleIdStr.trim());
} catch (NumberFormatException e) {
return null;
}
if (allowedRoleIds.isEmpty()) {
return roleId;
}
return allowedRoleIds.contains(roleId) ? roleId : null;
}
/**
* 解析分类隔离范围:角色有绑定时强制使用角色的分类,否则使用客户端传入的 categoryIds。
*/
private List<Long> resolveCategoryIds(RoleScope scope, String categoryIds) {
if (scope.hasRole()) {
return scope.categoryIds();
}
return categoryFilter.parse(categoryIds);
}
}