AI接入 v0.1
This commit is contained in:
@@ -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)接入配置。
|
||||
*
|
||||
* <p>走 OpenAI 兼容的 chat/completions 协议,**不绑定任何供应商** ——
|
||||
* DeepSeek / 通义(compatible-mode)/ 智谱 / OpenAI 都只需要改 base-url 和 model:
|
||||
* <pre>
|
||||
* 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
|
||||
* </pre>
|
||||
* 拼接规则是 {@code base-url + "/chat/completions"}。</p>
|
||||
*
|
||||
* <p>api-key 只从环境变量读,application.yml 里不写默认值。key 为空时
|
||||
* {@link #isReady()} 为 false,AI 接口自动降级为「未配置」提示,而不是 500。</p>
|
||||
*/
|
||||
@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。
|
||||
*
|
||||
* <p>刻意不复用现有的 RestTemplate —— 那个 5 秒读超时撑不住 LLM 的流式生成。
|
||||
* 用 JDK 自带的 HttpClient 是为了**零新增 Maven 依赖**,流式读响应体靠
|
||||
* {@code BodyHandlers.ofInputStream()}。</p>
|
||||
*/
|
||||
@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;
|
||||
});
|
||||
}
|
||||
}
|
||||
@@ -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 助手接口。
|
||||
*
|
||||
* <p>会话与消息都按 userId 隔离(和 note 等模块同一套纪律)。
|
||||
* /chat 走 SSE 流式,鉴权仍然是 JWT —— 没有在 SecurityConfig 里放行。</p>
|
||||
*/
|
||||
@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<String, Object> status() {
|
||||
return Map.of("enabled", aiConfig.isReady());
|
||||
}
|
||||
|
||||
@Operation(summary = "会话列表(按最近活跃倒序)")
|
||||
@GetMapping("/sessions")
|
||||
public List<SessionItem> sessions(Authentication authentication) {
|
||||
Long userId = getUserId(authentication);
|
||||
List<ChatSession> rows = sessionMapper.selectList(
|
||||
new LambdaQueryWrapper<ChatSession>()
|
||||
.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<ChatMessageItem> messages(@PathVariable Long id, Authentication authentication) {
|
||||
Long userId = getUserId(authentication);
|
||||
requireOwnedSession(userId, id);
|
||||
|
||||
return messageMapper.selectList(new LambdaQueryWrapper<ChatMessage>()
|
||||
.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<String, Object> deleteSession(@PathVariable Long id, Authentication authentication) {
|
||||
Long userId = getUserId(authentication);
|
||||
requireOwnedSession(userId, id);
|
||||
|
||||
sessionMapper.deleteById(id);
|
||||
// 会话没了,消息也一并隐藏 —— 不然跨会话搜索时会捞出孤儿消息
|
||||
messageMapper.update(null, new LambdaUpdateWrapper<ChatMessage>()
|
||||
.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<User>().eq(User::getUsername, username)
|
||||
);
|
||||
if (user == null) {
|
||||
throw new IllegalArgumentException("用户不存在");
|
||||
}
|
||||
return user.getId();
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
package com.accounting.dto.ai;
|
||||
|
||||
import jakarta.validation.constraints.NotBlank;
|
||||
import jakarta.validation.constraints.Size;
|
||||
import lombok.Data;
|
||||
|
||||
/**
|
||||
* 发起对话的请求。
|
||||
*
|
||||
* <p>只传会话 ID 和这条消息的内容 —— 历史由后端自己从 chat_message 表加载,
|
||||
* App 不用(也不该)把历史搬来搬去。sessionId 传 null 表示新开会话,
|
||||
* 后端会自动建一个、标题取这条消息的前 20 字,并在 done 事件里把新会话 ID 带回来。</p>
|
||||
*/
|
||||
@Data
|
||||
public class ChatSendRequest {
|
||||
|
||||
/** 会话 ID,null = 自动新建 */
|
||||
private Long sessionId;
|
||||
|
||||
@NotBlank(message = "消息内容不能为空")
|
||||
@Size(max = 4000, message = "单条消息最长 4000 字")
|
||||
private String content;
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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 消息
|
||||
*
|
||||
* <p>消息不可变。工具调用的中间过程(tool_calls / tool 结果)不入库 ——
|
||||
* 历史里只有 user / assistant 文本,App 端因此不需要理解工具消息的格式。</p>
|
||||
*/
|
||||
@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;
|
||||
}
|
||||
@@ -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 会话
|
||||
*
|
||||
* <p>deleted 字段与 application.yml 中的
|
||||
* {@code mybatis-plus.global-config.db-config.logic-delete-field=deleted} 对应,
|
||||
* 查询会自动追加 {@code deleted = 0},删除会自动变成 UPDATE。</p>
|
||||
*/
|
||||
@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;
|
||||
}
|
||||
@@ -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<ChatMessage> {
|
||||
}
|
||||
@@ -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<ChatSession> {
|
||||
}
|
||||
@@ -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 + 工具调用循环 + 消息落库。
|
||||
*
|
||||
* <p>链路:
|
||||
* <pre>
|
||||
* 入库 user 消息 → 加载历史 → 组装 system prompt
|
||||
* → [流式调 LLM → 有 tool_calls? 执行工具(脱敏) → 回填 → 再调] 最多 N 轮
|
||||
* → 入库 assistant 消息 → SSE done
|
||||
* </pre>
|
||||
* 两个刻意的设计:
|
||||
* <ul>
|
||||
* <li>**工具调用的中间过程不入库** —— 历史里只有 user/assistant 文本,
|
||||
* 所以发请求时可以放心地只从 chat_message 表重建上下文,App 端
|
||||
* 不需要理解 tool 消息格式。</li>
|
||||
* <li>**脱敏只发生在出站方向**(用户消息 + 工具结果),AI 回答里的
|
||||
* [REDACTED_x] 占位符不还原,见 {@link SensitiveDataMasker}。</li>
|
||||
* </ul>
|
||||
*/
|
||||
@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 线程池里进行。
|
||||
*
|
||||
* <p>统一约定:无论出什么错,客户端收到的都是 SSE 的 error 事件 ——
|
||||
* 这个接口永远不返回 4xx/5xx 的 JSON,App 端只要处理一种错误通道。</p>
|
||||
*/
|
||||
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<Map<String, Object>> 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<Map<String, Object>> 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<ToolCall> toolCalls, List<Object> rawToolCalls) {}
|
||||
|
||||
/** 流式调一次 chat/completions,边读边把文本增量推给 App */
|
||||
private Round streamOnce(List<Map<String, Object>> llmMessages, SseEmitter emitter) throws Exception {
|
||||
Map<String, Object> 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<java.io.InputStream> 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<Integer, String> ids = new TreeMap<>();
|
||||
Map<Integer, String> names = new TreeMap<>();
|
||||
Map<Integer, StringBuilder> 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<ToolCall> calls = new ArrayList<>();
|
||||
List<Object> 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<Map<String, Object>> buildLlmMessages(Long sessionId) {
|
||||
List<ChatMessage> history = messageMapper.selectList(
|
||||
new LambdaQueryWrapper<ChatMessage>()
|
||||
.eq(ChatMessage::getSessionId, sessionId)
|
||||
.orderByAsc(ChatMessage::getId));
|
||||
if (history.size() > aiConfig.getMaxHistoryMessages()) {
|
||||
history = history.subList(history.size() - aiConfig.getMaxHistoryMessages(), history.size());
|
||||
}
|
||||
|
||||
List<Map<String, Object>> 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);
|
||||
}
|
||||
}
|
||||
@@ -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 格式)。
|
||||
*
|
||||
* <p>v1 只暴露 3 个查询工具,**没有任何写操作** —— 让模型能改数据之前,
|
||||
* 必须先有「AI 生成建议 → 用户确认 → 落库」的确认流程,那是阶段 C 的事。</p>
|
||||
*
|
||||
* <p>工具抛出的业务异常不往外抛,而是包装成 {"error": ...} 返回给模型 ——
|
||||
* 让模型自己决定怎么向用户解释(比如「没有找到这篇笔记」),
|
||||
* 而不是让整次对话直接失败。</p>
|
||||
*/
|
||||
@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<Map<String, Object>> definitions() {
|
||||
List<Map<String, Object>> 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<Map<String, Object>> 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<String, Object> m = new LinkedHashMap<>();
|
||||
m.put("category", c.getCategoryName());
|
||||
m.put("amount", c.getAmount());
|
||||
return m;
|
||||
})
|
||||
.toList();
|
||||
|
||||
Map<String, Object> 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<Map<String, Object>> notes = new ArrayList<>();
|
||||
for (NoteListItem item : page.getList()) {
|
||||
Map<String, Object> m = new LinkedHashMap<>();
|
||||
m.put("id", item.getId());
|
||||
m.put("title", item.getTitle());
|
||||
m.put("summary", item.getSummary());
|
||||
notes.add(m);
|
||||
}
|
||||
|
||||
Map<String, Object> 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<String, Object> result = new LinkedHashMap<>();
|
||||
result.put("id", note.getId());
|
||||
result.put("title", note.getTitle());
|
||||
result.put("content", note.getContent());
|
||||
return objectMapper.writeValueAsString(result);
|
||||
}
|
||||
|
||||
private Map<String, Object> function(String name, String description, Map<String, Object> 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 + "\"}";
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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 的文本都必须先过这里。
|
||||
*
|
||||
* <p>为什么要做:用户的笔记里存着服务器 IP、SSH 端口、MySQL root 密码、
|
||||
* 私钥路径 —— 这些内容一旦被 AI 读取,就会原样发到第三方 API。</p>
|
||||
*
|
||||
* <p><b>v1 的一个刻意取舍:只做出站方向,AI 回答里的占位符不还原。</b>
|
||||
* 泄露风险在「发给第三方」那一段;用户自己手机上看到占位符反而更安全
|
||||
* (连截屏分享都不会泄露),代价是 AI 无法复述具体的密码/密钥 ——
|
||||
* 这是刻意的:敏感值本来就不该由 AI 复述。流式响应里做跨 chunk 的占位符
|
||||
* 还原既复杂又容易出 bug,不值得。</p>
|
||||
*
|
||||
* <p><b>要有预期</b>:正则脱敏是兜底不是银弹,总会有写法怪异漏网的。
|
||||
* 真正的第二道保险(按目录/标签配置 AI 可读范围)是 v2 的事。</p>
|
||||
*/
|
||||
@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 位以上、**同时含字母和数字**的连续串。
|
||||
*
|
||||
* <p>兜住 AK/SK(LTAI5tD...)、签名(32 位 hex)、JWT 这类没有键名前缀的裸密钥。</p>
|
||||
*
|
||||
* <p>「必须含字母和数字」这个条件不能省:少了它,`migration_add_note.sql`
|
||||
* 这种纯字母的长文件名/包名会被整个脱敏掉。(这个是单测抓出来的)</p>
|
||||
*/
|
||||
private static final Pattern LONG_TOKEN = Pattern.compile(
|
||||
"(?<![A-Za-z0-9+/=_.-])"
|
||||
+ "(?=[A-Za-z0-9+/=_.-]*\\d)(?=[A-Za-z0-9+/=_.-]*[A-Za-z])"
|
||||
+ "[A-Za-z0-9+/=_.-]{20,}"
|
||||
+ "(?![A-Za-z0-9+/=_.-])");
|
||||
|
||||
private static final Pattern IPV4 = Pattern.compile(
|
||||
"\\b(?:(?:25[0-5]|2[0-4]\\d|1\\d\\d|[1-9]?\\d)\\.){3}(?:25[0-5]|2[0-4]\\d|1\\d\\d|[1-9]?\\d)\\b");
|
||||
|
||||
/**
|
||||
* 脱敏。同一个原文在一次调用内映射到同一个占位符 ——
|
||||
* 这样模型看到的一致,回答不会前后对不上。
|
||||
*/
|
||||
public String mask(String text) {
|
||||
if (text == null || text.isEmpty()) {
|
||||
return text;
|
||||
}
|
||||
Map<String, String> 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<String, String> 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<String, String> 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<String, String> dict, int[] counter, String kind, String secret) {
|
||||
return dict.computeIfAbsent(secret, k -> "[REDACTED_" + kind + "_" + (++counter[0]) + "]");
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user