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]) + "]");
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -70,6 +70,20 @@ k780:
|
|||||||
server:
|
server:
|
||||||
port: 12345
|
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:
|
logging:
|
||||||
level:
|
level:
|
||||||
com.accounting: debug
|
com.accounting: debug
|
||||||
|
|||||||
@@ -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消息表';
|
||||||
@@ -179,3 +179,29 @@ INSERT INTO `category` (`user_id`, `name`, `icon`, `type`, `sort_order`) VALUES
|
|||||||
(NULL, '投资', '📈', 2, 3),
|
(NULL, '投资', '📈', 2, 3),
|
||||||
(NULL, '兼职', '💼', 2, 4),
|
(NULL, '兼职', '💼', 2, 4),
|
||||||
(NULL, '其他', '📦', 2, 5);
|
(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消息表';
|
||||||
|
|||||||
@@ -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;
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user