diff --git a/src/main/java/com/accounting/config/AiConfig.java b/src/main/java/com/accounting/config/AiConfig.java new file mode 100644 index 0000000..adabc18 --- /dev/null +++ b/src/main/java/com/accounting/config/AiConfig.java @@ -0,0 +1,128 @@ +package com.accounting.config; + +import jakarta.annotation.PostConstruct; +import lombok.extern.slf4j.Slf4j; +import org.springframework.beans.factory.annotation.Value; +import org.springframework.context.annotation.Bean; +import org.springframework.context.annotation.Configuration; + +import java.net.http.HttpClient; +import java.time.Duration; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; + +/** + * AI(LLM)接入配置。 + * + *

走 OpenAI 兼容的 chat/completions 协议,**不绑定任何供应商** —— + * DeepSeek / 通义(compatible-mode)/ 智谱 / OpenAI 都只需要改 base-url 和 model: + *

+ *   DeepSeek  https://api.deepseek.com/v1
+ *   通义      https://dashscope.aliyuncs.com/compatible-mode/v1
+ *   智谱      https://open.bigmodel.cn/api/paas/v4
+ *   OpenAI    https://api.openai.com/v1
+ * 
+ * 拼接规则是 {@code base-url + "/chat/completions"}。

+ * + *

api-key 只从环境变量读,application.yml 里不写默认值。key 为空时 + * {@link #isReady()} 为 false,AI 接口自动降级为「未配置」提示,而不是 500。

+ */ +@Slf4j +@Configuration +public class AiConfig { + + @Value("${ai.enabled:true}") + private boolean enabled; + + /** 不带 /chat/completions 后缀 */ + @Value("${ai.base-url:}") + private String baseUrl; + + @Value("${ai.api-key:}") + private String apiKey; + + @Value("${ai.model:}") + private String model; + + /** 工具调用最多循环几轮:模型连续要数据时防止无限循环烧 token */ + @Value("${ai.max-tool-rounds:3}") + private int maxToolRounds; + + /** 单次对话最多带多少条历史消息,太久远的不带 */ + @Value("${ai.max-history-messages:20}") + private int maxHistoryMessages; + + public boolean isReady() { + return enabled && !baseUrl.isBlank() && !apiKey.isBlank() && !model.isBlank(); + } + + /** + * 启动时把配置状态打到日志里 —— 配错了不用猜,看 app.log 就知道。 + * 注意**永远不打印 api-key 的内容**,只报长度。 + */ + @PostConstruct + public void logConfigStatus() { + if (isReady()) { + log.info("AI 已启用:base-url={}, model={}, api-key 长度={}", + getBaseUrl(), model, apiKey.length()); + } else { + log.warn("AI 未配置,助手功能将降级(status 返回 false,对话推 error 事件):" + + "enabled={}, base-url={}, model={}, api-key={}", + enabled, + baseUrl.isBlank() ? "(空)" : baseUrl, + model.isBlank() ? "(空)" : model, + apiKey.isBlank() ? "(空)" : "已设置"); + } + } + + /** + * 容忍结尾多余的斜杠:否则会拼出 {@code //chat/completions}, + * 网关直接 404,而且报错信息完全看不出是这个原因。 + */ + public String getBaseUrl() { + return baseUrl.replaceAll("/+$", ""); + } + + public String getApiKey() { + return apiKey; + } + + public String getModel() { + return model; + } + + public int getMaxToolRounds() { + return maxToolRounds; + } + + public int getMaxHistoryMessages() { + return maxHistoryMessages; + } + + /** + * 出站调 LLM 用的 HttpClient。 + * + *

刻意不复用现有的 RestTemplate —— 那个 5 秒读超时撑不住 LLM 的流式生成。 + * 用 JDK 自带的 HttpClient 是为了**零新增 Maven 依赖**,流式读响应体靠 + * {@code BodyHandlers.ofInputStream()}。

+ */ + @Bean + public HttpClient aiHttpClient() { + return HttpClient.newBuilder() + .connectTimeout(Duration.ofSeconds(10)) + .build(); + } + + /** + * SSE 的工作线程池:controller 立刻返回 SseEmitter, + * 上游的流式读取和工具循环都在这个池子里跑,不占用 Tomcat 请求线程。 + */ + @Bean(destroyMethod = "shutdown") + public ExecutorService aiExecutor() { + return Executors.newCachedThreadPool(r -> { + Thread t = new Thread(r, "ai-chat"); + t.setDaemon(true); + return t; + }); + } +} diff --git a/src/main/java/com/accounting/controller/AiController.java b/src/main/java/com/accounting/controller/AiController.java new file mode 100644 index 0000000..d874610 --- /dev/null +++ b/src/main/java/com/accounting/controller/AiController.java @@ -0,0 +1,152 @@ +package com.accounting.controller; + +import com.accounting.config.AiConfig; +import com.accounting.dto.ai.ChatMessageItem; +import com.accounting.dto.ai.ChatSendRequest; +import com.accounting.dto.ai.SessionItem; +import com.accounting.entity.ChatMessage; +import com.accounting.entity.ChatSession; +import com.accounting.entity.User; +import com.accounting.mapper.ChatMessageMapper; +import com.accounting.mapper.ChatSessionMapper; +import com.accounting.mapper.UserMapper; +import com.accounting.service.AiChatService; +import com.baomidou.mybatisplus.core.conditions.query.LambdaQueryWrapper; +import com.baomidou.mybatisplus.core.conditions.update.LambdaUpdateWrapper; +import io.swagger.v3.oas.annotations.Operation; +import io.swagger.v3.oas.annotations.tags.Tag; +import jakarta.validation.Valid; +import lombok.extern.slf4j.Slf4j; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.http.MediaType; +import org.springframework.security.core.Authentication; +import org.springframework.security.core.userdetails.UserDetails; +import org.springframework.web.bind.annotation.DeleteMapping; +import org.springframework.web.bind.annotation.GetMapping; +import org.springframework.web.bind.annotation.PathVariable; +import org.springframework.web.bind.annotation.PostMapping; +import org.springframework.web.bind.annotation.RequestBody; +import org.springframework.web.bind.annotation.RequestMapping; +import org.springframework.web.bind.annotation.RestController; +import org.springframework.web.servlet.mvc.method.annotation.SseEmitter; + +import java.util.List; +import java.util.Map; + +/** + * AI 助手接口。 + * + *

会话与消息都按 userId 隔离(和 note 等模块同一套纪律)。 + * /chat 走 SSE 流式,鉴权仍然是 JWT —— 没有在 SecurityConfig 里放行。

+ */ +@Slf4j +@Tag(name = "AI 助手") +@RestController +@RequestMapping("/api/ai") +public class AiController { + + /** 会话列表最多返回多少条 —— 单用户场景 100 条绰绰有余 */ + private static final int MAX_SESSIONS = 100; + + @Autowired + private AiConfig aiConfig; + + @Autowired + private AiChatService aiChatService; + + @Autowired + private ChatSessionMapper sessionMapper; + + @Autowired + private ChatMessageMapper messageMapper; + + @Autowired + private UserMapper userMapper; + + @Operation(summary = "AI 是否已配置(App 据此显示入口或提示)") + @GetMapping("/status") + public Map status() { + return Map.of("enabled", aiConfig.isReady()); + } + + @Operation(summary = "会话列表(按最近活跃倒序)") + @GetMapping("/sessions") + public List sessions(Authentication authentication) { + Long userId = getUserId(authentication); + List rows = sessionMapper.selectList( + new LambdaQueryWrapper() + .eq(ChatSession::getUserId, userId) + .orderByDesc(ChatSession::getUpdateTime) + .last("LIMIT " + MAX_SESSIONS)); + + return rows.stream().map(s -> { + SessionItem item = new SessionItem(); + item.setId(s.getId()); + item.setTitle(s.getTitle()); + item.setUpdateTime(s.getUpdateTime()); + return item; + }).toList(); + } + + @Operation(summary = "某个会话的消息(按时间正序)") + @GetMapping("/sessions/{id}/messages") + public List messages(@PathVariable Long id, Authentication authentication) { + Long userId = getUserId(authentication); + requireOwnedSession(userId, id); + + return messageMapper.selectList(new LambdaQueryWrapper() + .eq(ChatMessage::getSessionId, id) + .orderByAsc(ChatMessage::getId)) + .stream().map(m -> { + ChatMessageItem item = new ChatMessageItem(); + item.setId(m.getId()); + item.setRole(m.getRole()); + item.setContent(m.getContent()); + item.setCreateTime(m.getCreateTime()); + return item; + }).toList(); + } + + @Operation(summary = "删除会话(连消息一起逻辑删除)") + @DeleteMapping("/sessions/{id}") + public Map deleteSession(@PathVariable Long id, Authentication authentication) { + Long userId = getUserId(authentication); + requireOwnedSession(userId, id); + + sessionMapper.deleteById(id); + // 会话没了,消息也一并隐藏 —— 不然跨会话搜索时会捞出孤儿消息 + messageMapper.update(null, new LambdaUpdateWrapper() + .eq(ChatMessage::getSessionId, id) + .set(ChatMessage::getDeleted, 1)); + return Map.of("success", true); + } + + @Operation(summary = "发起对话(SSE 流式;sessionId 为空则自动新建会话)") + @PostMapping(value = "/chat", produces = MediaType.TEXT_EVENT_STREAM_VALUE) + public SseEmitter chat(@Valid @RequestBody ChatSendRequest request, Authentication authentication) { + Long userId = getUserId(authentication); + return aiChatService.chat(userId, request.getSessionId(), request.getContent()); + } + + // ---------------------------------------------------------------- 内部方法 + + private ChatSession requireOwnedSession(Long userId, Long sessionId) { + ChatSession session = sessionMapper.selectById(sessionId); + if (session == null || !session.getUserId().equals(userId)) { + throw new IllegalArgumentException("会话不存在"); + } + return session; + } + + private Long getUserId(Authentication authentication) { + UserDetails userDetails = (UserDetails) authentication.getPrincipal(); + String username = userDetails.getUsername(); + User user = userMapper.selectOne( + new LambdaQueryWrapper().eq(User::getUsername, username) + ); + if (user == null) { + throw new IllegalArgumentException("用户不存在"); + } + return user.getId(); + } +} diff --git a/src/main/java/com/accounting/dto/ai/ChatMessageItem.java b/src/main/java/com/accounting/dto/ai/ChatMessageItem.java new file mode 100644 index 0000000..566a465 --- /dev/null +++ b/src/main/java/com/accounting/dto/ai/ChatMessageItem.java @@ -0,0 +1,19 @@ +package com.accounting.dto.ai; + +import lombok.Data; + +import java.time.LocalDateTime; + +/** 会话内的一条消息(只有 user / assistant 文本,工具调用的中间过程不入库) */ +@Data +public class ChatMessageItem { + + private Long id; + + /** user / assistant */ + private String role; + + private String content; + + private LocalDateTime createTime; +} diff --git a/src/main/java/com/accounting/dto/ai/ChatSendRequest.java b/src/main/java/com/accounting/dto/ai/ChatSendRequest.java new file mode 100644 index 0000000..3c4d77d --- /dev/null +++ b/src/main/java/com/accounting/dto/ai/ChatSendRequest.java @@ -0,0 +1,23 @@ +package com.accounting.dto.ai; + +import jakarta.validation.constraints.NotBlank; +import jakarta.validation.constraints.Size; +import lombok.Data; + +/** + * 发起对话的请求。 + * + *

只传会话 ID 和这条消息的内容 —— 历史由后端自己从 chat_message 表加载, + * App 不用(也不该)把历史搬来搬去。sessionId 传 null 表示新开会话, + * 后端会自动建一个、标题取这条消息的前 20 字,并在 done 事件里把新会话 ID 带回来。

+ */ +@Data +public class ChatSendRequest { + + /** 会话 ID,null = 自动新建 */ + private Long sessionId; + + @NotBlank(message = "消息内容不能为空") + @Size(max = 4000, message = "单条消息最长 4000 字") + private String content; +} diff --git a/src/main/java/com/accounting/dto/ai/SessionItem.java b/src/main/java/com/accounting/dto/ai/SessionItem.java new file mode 100644 index 0000000..e911160 --- /dev/null +++ b/src/main/java/com/accounting/dto/ai/SessionItem.java @@ -0,0 +1,17 @@ +package com.accounting.dto.ai; + +import lombok.Data; + +import java.time.LocalDateTime; + +/** 会话列表项(不含消息) */ +@Data +public class SessionItem { + + private Long id; + + private String title; + + /** 最近活跃时间,列表按它倒序 */ + private LocalDateTime updateTime; +} diff --git a/src/main/java/com/accounting/entity/ChatMessage.java b/src/main/java/com/accounting/entity/ChatMessage.java new file mode 100644 index 0000000..de467f0 --- /dev/null +++ b/src/main/java/com/accounting/entity/ChatMessage.java @@ -0,0 +1,35 @@ +package com.accounting.entity; + +import com.baomidou.mybatisplus.annotation.IdType; +import com.baomidou.mybatisplus.annotation.TableId; +import com.baomidou.mybatisplus.annotation.TableLogic; +import com.baomidou.mybatisplus.annotation.TableName; +import lombok.Data; + +import java.time.LocalDateTime; + +/** + * AI 消息 + * + *

消息不可变。工具调用的中间过程(tool_calls / tool 结果)不入库 —— + * 历史里只有 user / assistant 文本,App 端因此不需要理解工具消息的格式。

+ */ +@Data +@TableName("chat_message") +public class ChatMessage { + + @TableId(type = IdType.AUTO) + private Long id; + + private Long sessionId; + + /** 角色:user / assistant */ + private String role; + + private String content; + + @TableLogic + private Integer deleted; + + private LocalDateTime createTime; +} diff --git a/src/main/java/com/accounting/entity/ChatSession.java b/src/main/java/com/accounting/entity/ChatSession.java new file mode 100644 index 0000000..aa5bca6 --- /dev/null +++ b/src/main/java/com/accounting/entity/ChatSession.java @@ -0,0 +1,37 @@ +package com.accounting.entity; + +import com.baomidou.mybatisplus.annotation.IdType; +import com.baomidou.mybatisplus.annotation.TableId; +import com.baomidou.mybatisplus.annotation.TableLogic; +import com.baomidou.mybatisplus.annotation.TableName; +import lombok.Data; + +import java.time.LocalDateTime; + +/** + * AI 会话 + * + *

deleted 字段与 application.yml 中的 + * {@code mybatis-plus.global-config.db-config.logic-delete-field=deleted} 对应, + * 查询会自动追加 {@code deleted = 0},删除会自动变成 UPDATE。

+ */ +@Data +@TableName("chat_session") +public class ChatSession { + + @TableId(type = IdType.AUTO) + private Long id; + + private Long userId; + + /** 标题,首条消息时自动取前 20 字 */ + private String title; + + @TableLogic + private Integer deleted; + + private LocalDateTime createTime; + + /** 最近活跃时间(DB 的 ON UPDATE 维护),会话列表按它倒序 */ + private LocalDateTime updateTime; +} diff --git a/src/main/java/com/accounting/mapper/ChatMessageMapper.java b/src/main/java/com/accounting/mapper/ChatMessageMapper.java new file mode 100644 index 0000000..14ccbf7 --- /dev/null +++ b/src/main/java/com/accounting/mapper/ChatMessageMapper.java @@ -0,0 +1,9 @@ +package com.accounting.mapper; + +import com.accounting.entity.ChatMessage; +import com.baomidou.mybatisplus.core.mapper.BaseMapper; +import org.apache.ibatis.annotations.Mapper; + +@Mapper +public interface ChatMessageMapper extends BaseMapper { +} diff --git a/src/main/java/com/accounting/mapper/ChatSessionMapper.java b/src/main/java/com/accounting/mapper/ChatSessionMapper.java new file mode 100644 index 0000000..4b2cb7e --- /dev/null +++ b/src/main/java/com/accounting/mapper/ChatSessionMapper.java @@ -0,0 +1,9 @@ +package com.accounting.mapper; + +import com.accounting.entity.ChatSession; +import com.baomidou.mybatisplus.core.mapper.BaseMapper; +import org.apache.ibatis.annotations.Mapper; + +@Mapper +public interface ChatSessionMapper extends BaseMapper { +} diff --git a/src/main/java/com/accounting/service/AiChatService.java b/src/main/java/com/accounting/service/AiChatService.java new file mode 100644 index 0000000..f8ab3a1 --- /dev/null +++ b/src/main/java/com/accounting/service/AiChatService.java @@ -0,0 +1,349 @@ +package com.accounting.service; + +import com.accounting.config.AiConfig; +import com.accounting.entity.ChatMessage; +import com.accounting.entity.ChatSession; +import com.accounting.mapper.ChatMessageMapper; +import com.accounting.mapper.ChatSessionMapper; +import com.baomidou.mybatisplus.core.conditions.query.LambdaQueryWrapper; +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; +import lombok.RequiredArgsConstructor; +import lombok.extern.slf4j.Slf4j; +import org.springframework.stereotype.Service; +import org.springframework.web.servlet.mvc.method.annotation.SseEmitter; + +import java.io.BufferedReader; +import java.io.InputStreamReader; +import java.net.URI; +import java.net.http.HttpClient; +import java.net.http.HttpRequest; +import java.net.http.HttpResponse; +import java.nio.charset.StandardCharsets; +import java.time.LocalDate; +import java.time.YearMonth; +import java.time.format.TextStyle; +import java.util.ArrayList; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Locale; +import java.util.Map; +import java.util.TreeMap; +import java.util.concurrent.ExecutorService; + +/** + * AI 对话编排:流式调 LLM + 工具调用循环 + 消息落库。 + * + *

链路: + *

+ *   入库 user 消息 → 加载历史 → 组装 system prompt
+ *     → [流式调 LLM → 有 tool_calls? 执行工具(脱敏) → 回填 → 再调] 最多 N 轮
+ *     → 入库 assistant 消息 → SSE done
+ * 
+ * 两个刻意的设计: + *
    + *
  • **工具调用的中间过程不入库** —— 历史里只有 user/assistant 文本, + * 所以发请求时可以放心地只从 chat_message 表重建上下文,App 端 + * 不需要理解 tool 消息格式。
  • + *
  • **脱敏只发生在出站方向**(用户消息 + 工具结果),AI 回答里的 + * [REDACTED_x] 占位符不还原,见 {@link SensitiveDataMasker}。
  • + *
+ */ +@Slf4j +@RequiredArgsConstructor +@Service +public class AiChatService { + + /** SSE 连接的最长寿命。LLM 生成慢 + 工具可能跑好几轮,给足余量 */ + private static final long EMITTER_TIMEOUT_MS = 180_000; + + private static final String ROLE_USER = "user"; + private static final String ROLE_ASSISTANT = "assistant"; + + private final AiConfig aiConfig; + private final HttpClient aiHttpClient; + private final ExecutorService aiExecutor; + private final ChatSessionMapper sessionMapper; + private final ChatMessageMapper messageMapper; + private final AiToolService toolService; + private final SensitiveDataMasker masker; + private final ObjectMapper objectMapper; + + /** + * 发起一次对话。立即返回 emitter,整个 LLM 交互在 AI 线程池里进行。 + * + *

统一约定:无论出什么错,客户端收到的都是 SSE 的 error 事件 —— + * 这个接口永远不返回 4xx/5xx 的 JSON,App 端只要处理一种错误通道。

+ */ + public SseEmitter chat(Long userId, Long sessionId, String content) { + SseEmitter emitter = new SseEmitter(EMITTER_TIMEOUT_MS); + emitter.onTimeout(emitter::complete); + emitter.onError(e -> log.warn("AI SSE 连接异常断开: {}", e.getMessage())); + + if (!aiConfig.isReady()) { + fail(emitter, "AI 功能未配置:请在服务端设置 AI_BASE_URL / AI_API_KEY / AI_MODEL"); + return emitter; + } + + aiExecutor.execute(() -> runChat(userId, sessionId, content, emitter)); + return emitter; + } + + // ---------------------------------------------------------------- 主流程 + + private void runChat(Long userId, Long sessionId, String content, SseEmitter emitter) { + try { + ChatSession session = resolveSession(userId, sessionId, content); + saveMessage(session.getId(), ROLE_USER, content); + + List> llmMessages = buildLlmMessages(session.getId()); + String assistantText = runCompletionLoop(userId, llmMessages, emitter); + + if (assistantText.isBlank()) { + throw new IllegalStateException("模型没有返回内容"); + } + saveMessage(session.getId(), ROLE_ASSISTANT, assistantText); + + send(emitter, "done", Map.of("sessionId", session.getId())); + emitter.complete(); + } catch (ClientGoneException e) { + log.info("AI 对话中止:客户端已断开"); + emitter.complete(); + } catch (Exception e) { + log.error("AI 对话失败", e); + fail(emitter, "AI 服务出错:" + e.getMessage()); + } + } + + /** + * 「流式调 LLM → 执行工具 → 回填 → 再调」的循环。 + * + * @return assistant 的最终文本(多轮的文本会拼在一起) + */ + private String runCompletionLoop(Long userId, List> llmMessages, + SseEmitter emitter) throws Exception { + StringBuilder full = new StringBuilder(); + int rounds = 0; + + while (true) { + Round round = streamOnce(llmMessages, emitter); + full.append(round.text()); + + if (round.toolCalls().isEmpty()) { + return full.toString(); + } + + rounds++; + if (rounds > aiConfig.getMaxToolRounds()) { + log.warn("AI 工具调用超过 {} 轮,强制截断", aiConfig.getMaxToolRounds()); + return full.toString().isBlank() ? "(连续查询次数超出上限,请换个问法)" : full.toString(); + } + + // 把 assistant 的 tool_calls 原样回填,再补上每个工具的结果 + llmMessages.add(Map.of( + "role", ROLE_ASSISTANT, + "content", round.text() == null ? "" : round.text(), + "tool_calls", round.rawToolCalls())); + + for (ToolCall call : round.toolCalls()) { + // 先告诉 App「正在查」,再执行 —— 查询可能要几百毫秒 + send(emitter, "tool", Map.of("name", call.name())); + String result = toolService.execute(userId, call.name(), call.arguments()); + llmMessages.add(Map.of( + "role", "tool", + "tool_call_id", call.id(), + "content", masker.mask(result))); + } + } + } + + // ---------------------------------------------------------------- 上游调用 + + private record ToolCall(String id, String name, String arguments) {} + + private record Round(String text, List toolCalls, List rawToolCalls) {} + + /** 流式调一次 chat/completions,边读边把文本增量推给 App */ + private Round streamOnce(List> llmMessages, SseEmitter emitter) throws Exception { + Map body = Map.of( + "model", aiConfig.getModel(), + "messages", llmMessages, + "tools", toolService.definitions(), + "stream", true); + + HttpRequest request = HttpRequest.newBuilder() + .uri(URI.create(aiConfig.getBaseUrl() + "/chat/completions")) + .timeout(java.time.Duration.ofSeconds(30)) // 等响应头的超时 + .header("Authorization", "Bearer " + aiConfig.getApiKey()) + .header("Content-Type", "application/json") + .header("Accept", "text/event-stream") + .POST(HttpRequest.BodyPublishers.ofString(objectMapper.writeValueAsString(body))) + .build(); + + HttpResponse response = + aiHttpClient.send(request, HttpResponse.BodyHandlers.ofInputStream()); + + if (response.statusCode() != 200) { + String err = new String(response.body().readAllBytes(), StandardCharsets.UTF_8); + throw new IllegalStateException("上游返回 " + response.statusCode() + ":" + abbreviate(err)); + } + + StringBuilder text = new StringBuilder(); + // tool_calls 的参数是分片到达的,按 index 累积 + Map ids = new TreeMap<>(); + Map names = new TreeMap<>(); + Map args = new LinkedHashMap<>(); + + try (BufferedReader reader = new BufferedReader( + new InputStreamReader(response.body(), StandardCharsets.UTF_8))) { + String line; + while ((line = reader.readLine()) != null) { + if (!line.startsWith("data:")) { + continue; + } + String data = line.substring(5).trim(); + if (data.equals("[DONE]")) { + break; + } + + JsonNode delta = objectMapper.readTree(data).path("choices").path(0).path("delta"); + + String content = delta.path("content").asText(null); + if (content != null && !content.isEmpty()) { + text.append(content); + send(emitter, "delta", Map.of("content", content)); + } + + JsonNode toolCalls = delta.path("tool_calls"); + if (toolCalls.isArray()) { + for (JsonNode tc : toolCalls) { + int idx = tc.path("index").asInt(0); + String id = tc.path("id").asText(null); + if (id != null && !id.isEmpty()) { + ids.put(idx, id); + } + String fn = tc.path("function").path("name").asText(null); + if (fn != null && !fn.isEmpty()) { + names.put(idx, fn); + } + args.computeIfAbsent(idx, k -> new StringBuilder()) + .append(tc.path("function").path("arguments").asText("")); + } + } + } + } + + List calls = new ArrayList<>(); + List raw = new ArrayList<>(); + for (Integer idx : ids.keySet()) { + String id = ids.get(idx); + String name = names.getOrDefault(idx, ""); + String arguments = args.getOrDefault(idx, new StringBuilder()).toString(); + calls.add(new ToolCall(id, name, arguments)); + raw.add(Map.of( + "id", id, + "type", "function", + "function", Map.of("name", name, "arguments", arguments))); + } + return new Round(text.toString(), calls, raw); + } + + // ---------------------------------------------------------------- 上下文组装 + + private List> buildLlmMessages(Long sessionId) { + List history = messageMapper.selectList( + new LambdaQueryWrapper() + .eq(ChatMessage::getSessionId, sessionId) + .orderByAsc(ChatMessage::getId)); + if (history.size() > aiConfig.getMaxHistoryMessages()) { + history = history.subList(history.size() - aiConfig.getMaxHistoryMessages(), history.size()); + } + + List> messages = new ArrayList<>(); + messages.add(Map.of("role", "system", "content", buildSystemPrompt())); + for (ChatMessage m : history) { + messages.add(Map.of("role", m.getRole(), "content", masker.mask(m.getContent()))); + } + return messages; + } + + private String buildSystemPrompt() { + LocalDate today = LocalDate.now(); + YearMonth thisMonth = YearMonth.from(today); + return """ + 你是「强宝小助手」的 AI 助手。这是一个个人助手应用,当前用户是它唯一的用户。\ + 你可以通过工具查询该用户的账单统计和笔记。 + + 今天的日期是 %s(%s)。用户说的「这个月」指 %s,「上个月」指 %s。 + + 回答要求: + - 直接给结论,简洁,不要客套。金额保留两位小数,前面加 ¥。 + - 涉及账单或笔记的具体数据时必须用工具查询,不要编造数字。 + - 如果工具结果里出现 [REDACTED_xxx] 这样的占位符,那是被有意脱敏的敏感信息,\ + 原样引用占位符即可,绝对不要猜测或编造它的内容。 + """.formatted( + today, + today.getDayOfWeek().getDisplayName(TextStyle.FULL, Locale.CHINA), + thisMonth, thisMonth.minusMonths(1)); + } + + // ---------------------------------------------------------------- 会话与消息 + + private ChatSession resolveSession(Long userId, Long sessionId, String firstContent) { + if (sessionId != null) { + ChatSession session = sessionMapper.selectById(sessionId); + if (session == null || !session.getUserId().equals(userId)) { + throw new IllegalArgumentException("会话不存在"); + } + return session; + } + ChatSession session = new ChatSession(); + session.setUserId(userId); + session.setTitle(abbreviate(firstContent.trim(), 20)); + sessionMapper.insert(session); + return session; + } + + private void saveMessage(Long sessionId, String role, String content) { + ChatMessage message = new ChatMessage(); + message.setSessionId(sessionId); + message.setRole(role); + message.setContent(content); + messageMapper.insert(message); + } + + // ---------------------------------------------------------------- SSE 工具方法 + + /** 客户端断开时把 send 的 IOException 转成这个,让上层安静退出而不是当错误处理 */ + private static class ClientGoneException extends RuntimeException { + } + + private void send(SseEmitter emitter, String event, Object payload) { + try { + emitter.send(SseEmitter.event().name(event).data(objectMapper.writeValueAsString(payload))); + } catch (Exception e) { + throw new ClientGoneException(); + } + } + + private void fail(SseEmitter emitter, String message) { + try { + send(emitter, "error", Map.of("message", message)); + emitter.complete(); + } catch (Exception ignored) { + // 连接已经没了,没什么可做的 + } + } + + /** 上游报错信息可能很长(一屏 HTML),截断后再往外抛 */ + private String abbreviate(String s) { + return abbreviate(s, 300); + } + + private String abbreviate(String s, int max) { + if (s == null) { + return ""; + } + return s.length() <= max ? s : s.substring(0, max); + } +} diff --git a/src/main/java/com/accounting/service/AiToolService.java b/src/main/java/com/accounting/service/AiToolService.java new file mode 100644 index 0000000..35579b1 --- /dev/null +++ b/src/main/java/com/accounting/service/AiToolService.java @@ -0,0 +1,180 @@ +package com.accounting.service; + +import com.accounting.dto.NoteListItem; +import com.accounting.dto.NotePageResponse; +import com.accounting.dto.NoteResponse; +import com.accounting.dto.StatisticsResponse; +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; +import lombok.RequiredArgsConstructor; +import lombok.extern.slf4j.Slf4j; +import org.springframework.stereotype.Component; + +import java.time.LocalDate; +import java.time.YearMonth; +import java.time.format.DateTimeParseException; +import java.util.ArrayList; +import java.util.Comparator; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; + +/** + * AI 的只读工具集(OpenAI function calling 格式)。 + * + *

v1 只暴露 3 个查询工具,**没有任何写操作** —— 让模型能改数据之前, + * 必须先有「AI 生成建议 → 用户确认 → 落库」的确认流程,那是阶段 C 的事。

+ * + *

工具抛出的业务异常不往外抛,而是包装成 {"error": ...} 返回给模型 —— + * 让模型自己决定怎么向用户解释(比如「没有找到这篇笔记」), + * 而不是让整次对话直接失败。

+ */ +@Slf4j +@RequiredArgsConstructor +@Component +public class AiToolService { + + private static final int MAX_NOTE_RESULTS = 10; + private static final int TOP_CATEGORIES = 5; + + private final StatisticsService statisticsService; + private final NoteService noteService; + private final ObjectMapper objectMapper; + + /** 工具定义(OpenAI tools 数组,直接塞进请求体) */ + public List> definitions() { + List> defs = new ArrayList<>(); + defs.add(function("query_bills", + "查询某个月份的收支汇总:总收入、总支出、结余,以及支出最多的前几个分类", + Map.of( + "type", "object", + "properties", Map.of( + "month", Map.of("type", "string", "description", "月份,格式 yyyy-MM,例如 2026-09")), + "required", List.of("month")))); + defs.add(function("query_notes", + "按关键词搜索用户的笔记,返回标题和摘要(不含正文)。要引用某篇笔记的内容时,先用它拿到笔记 id,再调 get_note", + Map.of( + "type", "object", + "properties", Map.of( + "keyword", Map.of("type", "string", "description", "关键词,匹配标题或正文,留空表示全部笔记")), + "required", List.of()))); + defs.add(function("get_note", + "读取一篇笔记的完整 Markdown 正文", + Map.of( + "type", "object", + "properties", Map.of( + "id", Map.of("type", "integer", "description", "笔记 ID,来自 query_notes 的结果")), + "required", List.of("id")))); + return defs; + } + + /** + * 执行一个工具调用。 + * + * @return 要回填给模型的内容(JSON 字符串,出站前会再过一遍脱敏) + */ + public String execute(Long userId, String name, String argumentsJson) { + log.info("AI tool call: user={}, tool={}, args={}", userId, name, argumentsJson); + try { + JsonNode args = objectMapper.readTree(argumentsJson == null || argumentsJson.isBlank() ? "{}" : argumentsJson); + return switch (name) { + case "query_bills" -> queryBills(userId, args); + case "query_notes" -> queryNotes(userId, args); + case "get_note" -> getNote(userId, args); + default -> errorResult("未知工具: " + name); + }; + } catch (IllegalArgumentException e) { + return errorResult(e.getMessage()); + } catch (Exception e) { + log.error("AI 工具执行失败: {}", name, e); + return errorResult("工具执行出错,请稍后重试"); + } + } + + private String queryBills(Long userId, JsonNode args) throws Exception { + String month = args.path("month").asText(""); + YearMonth ym; + try { + ym = YearMonth.parse(month); + } catch (DateTimeParseException e) { + return errorResult("month 格式不对,需要 yyyy-MM,例如 2026-09"); + } + + StatisticsResponse stats = statisticsService.getMonthlyStatistics( + userId, ym.getYear(), ym.getMonthValue()); + + // 只给模型支出方向、金额最大的前几个分类,省 token 也足够回答问题 + List> top = stats.getCategoryStatistics() == null ? List.of() + : stats.getCategoryStatistics().stream() + .filter(c -> c.getType() != null && c.getType() == 1) + .sorted(Comparator.comparing(StatisticsResponse.CategoryStatistics::getAmount, + Comparator.nullsLast(Comparator.reverseOrder()))) + .limit(TOP_CATEGORIES) + .map(c -> { + Map m = new LinkedHashMap<>(); + m.put("category", c.getCategoryName()); + m.put("amount", c.getAmount()); + return m; + }) + .toList(); + + Map result = new LinkedHashMap<>(); + result.put("month", month); + result.put("totalIncome", stats.getTotalIncome()); + result.put("totalExpense", stats.getTotalExpense()); + result.put("balance", stats.getBalance()); + result.put("topExpenseCategories", top); + return objectMapper.writeValueAsString(result); + } + + private String queryNotes(Long userId, JsonNode args) throws Exception { + String keyword = args.path("keyword").asText(""); + + NotePageResponse page = noteService.list(userId, "all", null, + keyword.isBlank() ? null : keyword, null, "updated", 1, MAX_NOTE_RESULTS); + + List> notes = new ArrayList<>(); + for (NoteListItem item : page.getList()) { + Map m = new LinkedHashMap<>(); + m.put("id", item.getId()); + m.put("title", item.getTitle()); + m.put("summary", item.getSummary()); + notes.add(m); + } + + Map result = new LinkedHashMap<>(); + result.put("total", page.getTotal()); + result.put("notes", notes); + return objectMapper.writeValueAsString(result); + } + + private String getNote(Long userId, JsonNode args) throws Exception { + if (!args.hasNonNull("id")) { + return errorResult("缺少笔记 id"); + } + NoteResponse note = noteService.detail(userId, args.get("id").asLong()); + + Map result = new LinkedHashMap<>(); + result.put("id", note.getId()); + result.put("title", note.getTitle()); + result.put("content", note.getContent()); + return objectMapper.writeValueAsString(result); + } + + private Map function(String name, String description, Map parameters) { + return Map.of( + "type", "function", + "function", Map.of( + "name", name, + "description", description, + "parameters", parameters)); + } + + private String errorResult(String message) { + try { + return objectMapper.writeValueAsString(Map.of("error", message)); + } catch (Exception e) { + return "{\"error\":\"" + message + "\"}"; + } + } +} diff --git a/src/main/java/com/accounting/service/SensitiveDataMasker.java b/src/main/java/com/accounting/service/SensitiveDataMasker.java new file mode 100644 index 0000000..5ab1699 --- /dev/null +++ b/src/main/java/com/accounting/service/SensitiveDataMasker.java @@ -0,0 +1,116 @@ +package com.accounting.service; + +import org.springframework.stereotype.Component; + +import java.util.HashMap; +import java.util.LinkedHashMap; +import java.util.Map; +import java.util.regex.Matcher; +import java.util.regex.Pattern; + +/** + * 出站脱敏:所有将要发给第三方 LLM API 的文本都必须先过这里。 + * + *

为什么要做:用户的笔记里存着服务器 IP、SSH 端口、MySQL root 密码、 + * 私钥路径 —— 这些内容一旦被 AI 读取,就会原样发到第三方 API。

+ * + *

v1 的一个刻意取舍:只做出站方向,AI 回答里的占位符不还原。 + * 泄露风险在「发给第三方」那一段;用户自己手机上看到占位符反而更安全 + * (连截屏分享都不会泄露),代价是 AI 无法复述具体的密码/密钥 —— + * 这是刻意的:敏感值本来就不该由 AI 复述。流式响应里做跨 chunk 的占位符 + * 还原既复杂又容易出 bug,不值得。

+ * + *

要有预期:正则脱敏是兜底不是银弹,总会有写法怪异漏网的。 + * 真正的第二道保险(按目录/标签配置 AI 可读范围)是 v2 的事。

+ */ +@Component +public class SensitiveDataMasker { + + /** 占位符形如 [REDACTED_IP_1]。模型会照抄它,所以形态要稳定 */ + public static final String PLACEHOLDER = "[REDACTED_"; + + /** PEM 私钥整块(多行)。必须最先处理,块内可能还嵌着其他形态 */ + private static final Pattern PRIVATE_KEY_BLOCK = Pattern.compile( + "-----BEGIN [A-Z ]*PRIVATE KEY-----[\\s\\S]*?-----END [A-Z ]*PRIVATE KEY-----"); + + /** password=xxx / 密码:xxx / passwd: xxx。group(1)=键名+分隔符(保留),group(2)=值(脱敏)。值 4 位起,避免把「密码:无」这种误伤 */ + private static final Pattern PASSWORD_KV = Pattern.compile( + "(?i)((?:password|passwd|pwd|密码|口令)\\s*[:=:]\\s*)(\"[^\"]{1,80}\"|\\S{4,80})"); + + /** api key / secret / token / sign 等键值形态。值限定 6+ 位,避免把 `token: 无` 这种误伤 */ + private static final Pattern SECRET_KV = Pattern.compile( + "(?i)\\b((?:api[_-]?key|access[_-]?key(?:[_-]?(?:id|secret))?|secret(?:[_-]?key)?" + + "|app[_-]?(?:key|secret)|token|sign)\\s*[:=:]\\s*)" + + "(\"[^\"]{6,120}\"|\\S{6,120})"); + + /** + * 长随机串:20 位以上、**同时含字母和数字**的连续串。 + * + *

兜住 AK/SK(LTAI5tD...)、签名(32 位 hex)、JWT 这类没有键名前缀的裸密钥。

+ * + *

「必须含字母和数字」这个条件不能省:少了它,`migration_add_note.sql` + * 这种纯字母的长文件名/包名会被整个脱敏掉。(这个是单测抓出来的)

+ */ + private static final Pattern LONG_TOKEN = Pattern.compile( + "(? dict = new HashMap<>(); + int[] counter = {0}; + + String out = replaceWhole(text, PRIVATE_KEY_BLOCK, "KEY", dict, counter); + out = replaceValueOnly(out, PASSWORD_KV, "SECRET", dict, counter); + out = replaceValueOnly(out, SECRET_KV, "SECRET", dict, counter); + out = replaceWhole(out, LONG_TOKEN, "TOKEN", dict, counter); + out = replaceWhole(out, IPV4, "IP", dict, counter); + return out; + } + + /** 整个匹配替换成占位符 */ + private String replaceWhole(String input, Pattern pattern, String kind, + Map dict, int[] counter) { + Matcher m = pattern.matcher(input); + if (!m.find()) { + return input; + } + StringBuffer sb = new StringBuffer(); + do { + m.appendReplacement(sb, Matcher.quoteReplacement(placeholder(dict, counter, kind, m.group()))); + } while (m.find()); + m.appendTail(sb); + return sb.toString(); + } + + /** 只把 group(2)(值)替换成占位符,group(1)(键名)保留 —— 保留结构有助于模型理解上下文 */ + private String replaceValueOnly(String input, Pattern pattern, String kind, + Map dict, int[] counter) { + Matcher m = pattern.matcher(input); + if (!m.find()) { + return input; + } + StringBuffer sb = new StringBuffer(); + do { + String replacement = m.group(1) + placeholder(dict, counter, kind, m.group(2)); + m.appendReplacement(sb, Matcher.quoteReplacement(replacement)); + } while (m.find()); + m.appendTail(sb); + return sb.toString(); + } + + private String placeholder(Map dict, int[] counter, String kind, String secret) { + return dict.computeIfAbsent(secret, k -> "[REDACTED_" + kind + "_" + (++counter[0]) + "]"); + } +} diff --git a/src/main/resources/application.yml b/src/main/resources/application.yml index 6a9c17a..9e9902f 100644 --- a/src/main/resources/application.yml +++ b/src/main/resources/application.yml @@ -70,6 +70,20 @@ k780: server: port: 12345 +# AI(LLM)配置:OpenAI 兼容协议,不绑定供应商 —— 换 DeepSeek/通义/智谱/OpenAI 只改环境变量 +ai: + enabled: true + # 拼接规则:base-url + /chat/completions,常用取值: + # DeepSeek https://api.deepseek.com/v1 + # 通义 https://dashscope.aliyuncs.com/compatible-mode/v1 + # 智谱 https://open.bigmodel.cn/api/paas/v4 + # OpenAI https://api.openai.com/v1 + base-url: ${AI_BASE_URL:} + api-key: ${AI_API_KEY:} # 只走环境变量,不留默认值;为空时 AI 功能自动降级为「未配置」 + model: ${AI_MODEL:} + max-tool-rounds: 3 # 工具调用最多循环几轮,防失控烧 token + max-history-messages: 20 # 单次对话最多带多少条历史 + logging: level: com.accounting: debug diff --git a/src/main/resources/db/migration_add_chat.sql b/src/main/resources/db/migration_add_chat.sql new file mode 100644 index 0000000..ad8f826 --- /dev/null +++ b/src/main/resources/db/migration_add_chat.sql @@ -0,0 +1,30 @@ +-- AI 会话与消息表(2026-09-23 阶段 B:AI 基建) +-- 设计说明: +-- * 与项目其他表一致:utf8mb4_unicode_ci、逻辑删除、指向 user 表的外键 +-- * 消息不可变(没有 update_time)——工具调用的中间过程不入库, +-- 历史里只有 user / assistant 文本,这样 App 端不用理解 tool 消息格式 +-- * 会话的 update_time 由 DB 的 ON UPDATE 维护,列表按「最近活跃」排序 + +CREATE TABLE IF NOT EXISTS `chat_session` ( + `id` BIGINT NOT NULL AUTO_INCREMENT COMMENT '会话ID', + `user_id` BIGINT NOT NULL COMMENT '用户ID', + `title` VARCHAR(100) NOT NULL DEFAULT '新对话' COMMENT '标题(首条消息自动生成)', + `deleted` TINYINT NOT NULL DEFAULT 0 COMMENT '逻辑删除:0正常 1已删除', + `create_time` DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP COMMENT '创建时间', + `update_time` DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP COMMENT '最近活跃时间', + PRIMARY KEY (`id`), + INDEX `idx_user_updated` (`user_id`, `deleted`, `update_time`), + FOREIGN KEY (`user_id`) REFERENCES `user` (`id`) ON DELETE CASCADE +) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci COMMENT='AI会话表'; + +CREATE TABLE IF NOT EXISTS `chat_message` ( + `id` BIGINT NOT NULL AUTO_INCREMENT COMMENT '消息ID', + `session_id` BIGINT NOT NULL COMMENT '会话ID', + `role` VARCHAR(16) NOT NULL COMMENT '角色:user / assistant', + `content` TEXT NOT NULL COMMENT '消息文本', + `deleted` TINYINT NOT NULL DEFAULT 0 COMMENT '逻辑删除:0正常 1已删除', + `create_time` DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP COMMENT '创建时间', + PRIMARY KEY (`id`), + INDEX `idx_session_created` (`session_id`, `deleted`, `create_time`), + FOREIGN KEY (`session_id`) REFERENCES `chat_session` (`id`) ON DELETE CASCADE +) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci COMMENT='AI消息表'; diff --git a/src/main/resources/db/schema.sql b/src/main/resources/db/schema.sql index 93255b5..8519ec6 100644 --- a/src/main/resources/db/schema.sql +++ b/src/main/resources/db/schema.sql @@ -179,3 +179,29 @@ INSERT INTO `category` (`user_id`, `name`, `icon`, `type`, `sort_order`) VALUES (NULL, '投资', '📈', 2, 3), (NULL, '兼职', '💼', 2, 4), (NULL, '其他', '📦', 2, 5); + +-- AI 会话表 +CREATE TABLE IF NOT EXISTS `chat_session` ( + `id` BIGINT NOT NULL AUTO_INCREMENT COMMENT '会话ID', + `user_id` BIGINT NOT NULL COMMENT '用户ID', + `title` VARCHAR(100) NOT NULL DEFAULT '新对话' COMMENT '标题(首条消息自动生成)', + `deleted` TINYINT NOT NULL DEFAULT 0 COMMENT '逻辑删除:0正常 1已删除', + `create_time` DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP COMMENT '创建时间', + `update_time` DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP COMMENT '最近活跃时间', + PRIMARY KEY (`id`), + INDEX `idx_user_updated` (`user_id`, `deleted`, `update_time`), + FOREIGN KEY (`user_id`) REFERENCES `user` (`id`) ON DELETE CASCADE +) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci COMMENT='AI会话表'; + +-- AI 消息表(工具调用的中间过程不入库,历史里只有 user / assistant 文本) +CREATE TABLE IF NOT EXISTS `chat_message` ( + `id` BIGINT NOT NULL AUTO_INCREMENT COMMENT '消息ID', + `session_id` BIGINT NOT NULL COMMENT '会话ID', + `role` VARCHAR(16) NOT NULL COMMENT '角色:user / assistant', + `content` TEXT NOT NULL COMMENT '消息文本', + `deleted` TINYINT NOT NULL DEFAULT 0 COMMENT '逻辑删除:0正常 1已删除', + `create_time` DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP COMMENT '创建时间', + PRIMARY KEY (`id`), + INDEX `idx_session_created` (`session_id`, `deleted`, `create_time`), + FOREIGN KEY (`session_id`) REFERENCES `chat_session` (`id`) ON DELETE CASCADE +) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci COMMENT='AI消息表'; diff --git a/src/test/java/com/accounting/service/SensitiveDataMaskerTest.java b/src/test/java/com/accounting/service/SensitiveDataMaskerTest.java new file mode 100644 index 0000000..f03109b --- /dev/null +++ b/src/test/java/com/accounting/service/SensitiveDataMaskerTest.java @@ -0,0 +1,100 @@ +package com.accounting.service; + +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertTrue; + +/** + * 脱敏器的单测 —— 这是整个 AI 链路里最不能出错的一个类, + * 用例按「用户笔记里真实出现过的形态」来写。 + */ +class SensitiveDataMaskerTest { + + private final SensitiveDataMasker masker = new SensitiveDataMasker(); + + @Test + @DisplayName("IPv4 地址被脱敏") + void masksIpv4() { + String out = masker.mask("服务器 IP 是 155.103.159.109,SSH 端口 31022"); + assertFalse(out.contains("155.103.159.109")); + assertTrue(out.contains("[REDACTED_IP_")); + assertTrue(out.contains("31022")); // 端口不算敏感,保留 + } + + @Test + @DisplayName("password 键值对:键名保留,值被替换") + void masksPasswordKv() { + String out = masker.mask("MySQL root 密码:mysql_jP65Fc,登录后记得改"); + assertFalse(out.contains("mysql_jP65Fc")); + assertTrue(out.contains("密码:[REDACTED_SECRET_")); + } + + @Test + @DisplayName("password= 等号形式同样生效") + void masksPasswordWithEquals() { + String out = masker.mask("jdbc:url?user=root&password=mySecret123"); + assertFalse(out.contains("mySecret123")); + } + + @Test + @DisplayName("PEM 私钥整块被脱敏") + void masksPrivateKeyBlock() { + String key = "-----BEGIN OPENSSH PRIVATE KEY-----\n" + + "b3BlbnNzaC1rZXktdjEAAAAA\n" + + "b3BlbnNzaC1rZXktdjEAAAAA\n" + + "-----END OPENSSH PRIVATE KEY-----"; + String out = masker.mask("私钥如下:\n" + key + "\n记得保管好"); + assertFalse(out.contains("b3BlbnNzaC1rZXktdjEAAAAA")); + assertTrue(out.contains("[REDACTED_KEY_1]")); + } + + @Test + @DisplayName("无键名的 AK/SK 长随机串被脱敏") + void masksBareAccessToken() { + String out = masker.mask("access key 是 LTAI5tDCJuB9YgLx4KeJwc9C 备用"); + assertFalse(out.contains("LTAI5tDCJuB9YgLx4KeJwc9C")); + } + + @Test + @DisplayName("同一个值在多次出现时映射到同一个占位符") + void keepsPlaceholderConsistent() { + String out = masker.mask("密码:abc12345,重复一下密码:abc12345"); + String first = out.substring(out.indexOf("[REDACTED_")); + assertEquals(1, countOccurrences(out, "[REDACTED_SECRET_1]"), "同一值应得到同一占位符: " + out); + assertFalse(out.contains(first + "[")); + } + + @Test + @DisplayName("普通内容不误伤:中文、短数字、文件名都保留") + void keepsNormalContent() { + String text = "笔记在 migration_add_note.sql,共 3 篇,端口 31022,金额 ¥134.00"; + assertEquals(text, masker.mask(text)); + } + + @Test + @DisplayName("「密码:无」这种非秘密不被误伤") + void keepsShortNonSecretValues() { + String out = masker.mask("密码:无"); + assertTrue(out.contains("密码:无")); + } + + @Test + @DisplayName("null 和空串原样返回") + void handlesNullAndEmpty() { + assertEquals(null, masker.mask(null)); + assertEquals("", masker.mask("")); + } + + private int countOccurrences(String text, String target) { + int count = 0; + int idx = 0; + while ((idx = text.indexOf(target, idx)) >= 0) { + count++; + idx += target.length(); + } + return count; + } +}