AI接入 v0.1

This commit is contained in:
2026-09-23 17:16:40 +08:00
parent ec21ba0876
commit be3fa61b8d
16 changed files with 1244 additions and 0 deletions
@@ -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]) + "]");
}
}
+14
View File
@@ -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
@@ -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消息表';
+26
View File
@@ -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消息表';
@@ -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;
}
}