Browse Source

提交代码

liuchengsen 1 tháng trước cách đây
mục cha
commit
9e5b36d332

+ 48 - 21
backend-java/src/main/java/com/pharmacopoeia/controller/ChatController.java

@@ -6,9 +6,10 @@ import com.pharmacopoeia.dto.FeedbackRequest;
 import com.pharmacopoeia.dto.ImageChatRequest;
 import com.pharmacopoeia.dto.MultimodalChatRequest;
 import com.pharmacopoeia.service.*;
-import org.springframework.http.MediaType;
+import jakarta.servlet.http.HttpServletRequest;
 import org.springframework.http.ResponseEntity;
 import org.springframework.http.codec.ServerSentEvent;
+import org.springframework.security.core.context.SecurityContextHolder;
 import org.springframework.web.bind.annotation.*;
 import reactor.core.publisher.Flux;
 import reactor.core.publisher.Sinks;
@@ -33,11 +34,13 @@ public class ChatController {
     private final QACacheService qaCache;
     private final JdbcTemplate jdbc;
     private final QwenProperties props;
+    private final HttpServletRequest request;
 
     public ChatController(RetrieverService rs, LLMService ls, PromptService ps,
                           ChatPersistenceService cps, RerankerService rrs,
                           QACacheService qaCache,
-                          JdbcTemplate jdbc, QwenProperties props) {
+                          JdbcTemplate jdbc, QwenProperties props,
+                          HttpServletRequest request) {
         this.retrieverService = rs;
         this.llmService = ls;
         this.promptService = ps;
@@ -46,6 +49,7 @@ public class ChatController {
         this.qaCache = qaCache;
         this.jdbc = jdbc;
         this.props = props;
+        this.request = request;
     }
 
     @PostMapping("/ask")
@@ -63,8 +67,8 @@ public class ChatController {
             String intent = (String) cached.getOrDefault("intent", "");
             @SuppressWarnings("unchecked")
             List<Map<String, Object>> sources = (List<Map<String, Object>>) cached.getOrDefault("sources", List.of());
-            persistenceService.saveMessage(cid, "user", query, intent, null);
-            persistenceService.saveMessage(cid, "assistant", answer, intent, sources);
+            persistenceService.saveMessage(getCurrentUserKey(), cid, "user", query, intent, null);
+            persistenceService.saveMessage(getCurrentUserKey(), cid, "assistant", answer, intent, sources);
             return ResponseEntity.ok(Map.of(
                     "answer", answer,
                     "sources", sources,
@@ -82,8 +86,8 @@ public class ChatController {
                 String intent = (String) waited.getOrDefault("intent", "");
                 @SuppressWarnings("unchecked")
                 List<Map<String, Object>> sources = (List<Map<String, Object>>) waited.getOrDefault("sources", List.of());
-                persistenceService.saveMessage(cid, "user", query, intent, null);
-                persistenceService.saveMessage(cid, "assistant", answer, intent, sources);
+                persistenceService.saveMessage(getCurrentUserKey(), cid, "user", query, intent, null);
+                persistenceService.saveMessage(getCurrentUserKey(), cid, "assistant", answer, intent, sources);
                 return ResponseEntity.ok(Map.of(
                         "answer", answer, "sources", sources, "intent", intent,
                         "conversation_id", cid, "cached", true
@@ -101,8 +105,8 @@ public class ChatController {
         String answer = llmAnswer;
 
         List<Map<String, Object>> sources = buildSources(docs);
-        persistenceService.saveMessage(cid, "user", query, intent, null);
-        persistenceService.saveMessage(cid, "assistant", answer, intent, sources);
+        persistenceService.saveMessage(getCurrentUserKey(), cid, "user", query, intent, null);
+        persistenceService.saveMessage(getCurrentUserKey(), cid, "assistant", answer, intent, sources);
 
         // 写入全局缓存
         qaCache.put(normalized, answer, intent, sources);
@@ -168,8 +172,8 @@ public class ChatController {
                                 } catch (Exception ignored) {}
 
                                 String finalAnswer = cleanAnswer(fullAnswer.toString());
-                                persistenceService.saveMessage(cid, "user", query, intent, null);
-                                persistenceService.saveMessage(cid, "assistant", finalAnswer, intent, sources);
+                                persistenceService.saveMessage(getCurrentUserKey(), cid, "user", query, intent, null);
+                                persistenceService.saveMessage(getCurrentUserKey(), cid, "assistant", finalAnswer, intent, sources);
 
                                 // 写入全局缓存
                                 qaCache.put(normalized, finalAnswer, intent, sources);
@@ -191,8 +195,8 @@ public class ChatController {
         @SuppressWarnings("unchecked")
         List<Map<String, Object>> sources = (List<Map<String, Object>>) cached.getOrDefault("sources", List.of());
 
-        persistenceService.saveMessage(cid, "user", query, intent, null);
-        persistenceService.saveMessage(cid, "assistant", answer, intent, sources);
+        persistenceService.saveMessage(getCurrentUserKey(), cid, "user", query, intent, null);
+        persistenceService.saveMessage(getCurrentUserKey(), cid, "assistant", answer, intent, sources);
 
         return Flux.create(sink -> {
             sink.next(ServerSentEvent.<String>builder().event("intent").data(intent).build());
@@ -245,10 +249,10 @@ public class ChatController {
         String answer = llmAnswer;
 
         List<Map<String, Object>> sources = buildSources(docs);
-        persistenceService.saveMessage(cid, "user",
+        persistenceService.saveMessage(getCurrentUserKey(), cid, "user",
                 request.getMessage().isBlank() ? "[图片]" : "[图片] " + request.getMessage(),
                 intent, null);
-        persistenceService.saveMessage(cid, "assistant", answer, intent, sources);
+        persistenceService.saveMessage(getCurrentUserKey(), cid, "assistant", answer, intent, sources);
 
         return ResponseEntity.ok(Map.of(
                 "answer", answer,
@@ -322,10 +326,10 @@ public class ChatController {
                                             } catch (Exception ignored) {}
 
                                             String finalAnswer = cleanAnswer(fullAnswer.toString());
-                                            persistenceService.saveMessage(cid, "user",
+                                            persistenceService.saveMessage(getCurrentUserKey(), cid, "user",
                                                     request.getMessage().isBlank() ? "[图片]" : "[图片] " + request.getMessage(),
                                                     intent, null);
-                                            persistenceService.saveMessage(cid, "assistant", finalAnswer, intent, sources);
+                                            persistenceService.saveMessage(getCurrentUserKey(), cid, "assistant", finalAnswer, intent, sources);
                                             sink.tryEmitComplete();
                                         })
                                         .doOnError(sink::tryEmitError)
@@ -390,8 +394,8 @@ public class ChatController {
         List<Map<String, Object>> sources = buildSources(docs);
         String userMsg = !request.getMessage().isBlank() ? request.getMessage()
                 : !ocrText.isEmpty() ? "[" + mediaLabel + "]" : request.getMessage();
-        persistenceService.saveMessage(cid, "user", userMsg, intent, null);
-        persistenceService.saveMessage(cid, "assistant", answer, intent, sources);
+        persistenceService.saveMessage(getCurrentUserKey(), cid, "user", userMsg, intent, null);
+        persistenceService.saveMessage(getCurrentUserKey(), cid, "assistant", answer, intent, sources);
 
         return ResponseEntity.ok(Map.of(
                 "answer", answer,
@@ -585,8 +589,8 @@ public class ChatController {
                                 String finalAnswer = cleanAnswer(fullAnswer.toString());
                                 String userMsg = !request.getMessage().isBlank() ? request.getMessage()
                                         : !ocrText.isEmpty() ? "[" + mediaLabel + "]" : "";
-                                persistenceService.saveMessage(cid, "user", userMsg, intent, null);
-                                persistenceService.saveMessage(cid, "assistant", finalAnswer, intent, sources);
+                                persistenceService.saveMessage(getCurrentUserKey(), cid, "user", userMsg, intent, null);
+                                persistenceService.saveMessage(getCurrentUserKey(), cid, "assistant", finalAnswer, intent, sources);
                                 sink.tryEmitComplete();
                             })
                             .doOnError(sink::tryEmitError)
@@ -599,7 +603,7 @@ public class ChatController {
             @RequestParam(defaultValue = "1") int page,
             @RequestParam(defaultValue = "20") int pageSize) {
         // getHistory 现在直接返回包含 items/page/page_size/total/total_pages 的 Map
-        var result = persistenceService.getHistory(page, pageSize);
+        var result = persistenceService.getHistory(getCurrentUserKey(), page, pageSize);
         return ResponseEntity.ok(result);
     }
 
@@ -609,6 +613,14 @@ public class ChatController {
         return ResponseEntity.ok(Map.of("conversation_id", cid, "messages", msgs));
     }
 
+    /** 返回最近 N 条消息,供前端恢复对话(微信 WebView 等 IndexedDB 不可用场景) */
+    @GetMapping("/recent-messages")
+    public ResponseEntity<Map<String, Object>> getRecentMessages(
+            @RequestParam(defaultValue = "50") int limit) {
+        var msgs = persistenceService.getRecentMessages(getCurrentUserKey(), Math.min(limit, 200));
+        return ResponseEntity.ok(Map.of("messages", msgs));
+    }
+
     @PostMapping("/feedback")
     public ResponseEntity<Map<String, Object>> submitFeedback(@RequestBody FeedbackRequest request) {
         persistenceService.updateFeedback(request.getMessageId(), request.getFeedback());
@@ -751,6 +763,21 @@ public class ChatController {
         return storedSection;
     }
 
+    /** 从 SecurityContext 获取当前用户标识(JWT subject),未登录则用 IP 隔离 */
+    private String getCurrentUserKey() {
+        var auth = SecurityContextHolder.getContext().getAuthentication();
+        if (auth != null && auth.isAuthenticated() && !"anonymousUser".equals(auth.getPrincipal())) {
+            return auth.getName();
+        }
+        // 未登录用 IP 隔离,避免不同手机会话串了
+        String ip = request.getRemoteAddr();
+        String forwarded = request.getHeader("X-Forwarded-For");
+        if (forwarded != null && !forwarded.isBlank()) {
+            ip = forwarded.split(",")[0].trim();
+        }
+        return "ip:" + ip;
+    }
+
     private String cleanAnswer(String text) {
         if (text == null) {
             return "";

+ 3 - 0
backend-java/src/main/java/com/pharmacopoeia/entity/Conversation.java

@@ -20,6 +20,9 @@ public class Conversation {
     @Column(name = "conversation_id", unique = true, nullable = false, length = 64)
     private String conversationId;
 
+    @Column(name = "user_key", length = 128)
+    private String userKey;
+
     @Column(name = "user_id")
     private Long userId;
 

+ 3 - 0
backend-java/src/main/java/com/pharmacopoeia/entity/Message.java

@@ -27,6 +27,9 @@ public class Message {
     @Column(name = "conversation_id", nullable = false, length = 64)
     private String conversationId;
 
+    @Column(name = "user_key", length = 128)
+    private String userKey;
+
     @Column(nullable = false, length = 32)
     private String role;
 

+ 3 - 5
backend-java/src/main/java/com/pharmacopoeia/repository/ConversationRepository.java

@@ -12,8 +12,6 @@ public interface ConversationRepository extends JpaRepository<Conversation, Long
     Optional<Conversation> findByConversationId(String conversationId);
     List<Conversation> findByUserIdOrderByCreatedAtDesc(Long userId);
 
-    /**
-     * 按创建时间倒序分页查询,避免全量加载到内存导致 OOM。
-     */
-    Page<Conversation> findAllByOrderByCreatedAtDesc(Pageable pageable);
-}
+    /** 按用户分页查询对话列表 */
+    Page<Conversation> findByUserKeyOrderByCreatedAtDesc(String userKey, Pageable pageable);
+}

+ 7 - 5
backend-java/src/main/java/com/pharmacopoeia/repository/MessageRepository.java

@@ -1,6 +1,8 @@
 package com.pharmacopoeia.repository;
 
 import com.pharmacopoeia.entity.Message;
+import org.springframework.data.domain.Page;
+import org.springframework.data.domain.Pageable;
 import org.springframework.data.jpa.repository.JpaRepository;
 import org.springframework.data.jpa.repository.Query;
 import org.springframework.data.repository.query.Param;
@@ -12,10 +14,10 @@ public interface MessageRepository extends JpaRepository<Message, Long> {
     List<Message> findByConversationIdOrderByCreatedAtAsc(String conversationId);
     long countByConversationId(String conversationId);
 
-    /**
-     * 按 conversationId 列表批量统计消息数,解决 N+1 查询问题。
-     * 返回每行 [conversationId(String), count(Long)]。
-     */
     @Query("SELECT m.conversationId, COUNT(m) FROM Message m WHERE m.conversationId IN :cids GROUP BY m.conversationId")
     List<Object[]> countByConversationIds(@Param("cids") List<String> cids);
-}
+
+    /** 按用户获取最近消息(倒序),用于前端恢复对话 */
+    @Query("SELECT m FROM Message m WHERE m.userKey = :userKey ORDER BY m.createdAt DESC")
+    Page<Message> findRecentByUserKey(@Param("userKey") String userKey, Pageable pageable);
+}

+ 23 - 13
backend-java/src/main/java/com/pharmacopoeia/service/ChatPersistenceService.java

@@ -28,11 +28,12 @@ public class ChatPersistenceService {
     }
 
     @Transactional
-    public void saveMessage(String conversationId, String role, String content,
+    public void saveMessage(String userKey, String conversationId, String role, String content,
                             String intent, List<Map<String, Object>> sources) {
         conversationRepository.findByConversationId(conversationId).orElseGet(() -> {
             var conv = Conversation.builder()
                     .conversationId(conversationId)
+                    .userKey(userKey)
                     .userId(0L)
                     .title(role.equals("user") && content.length() > 0
                             ? content.substring(0, Math.min(50, content.length()))
@@ -43,6 +44,7 @@ public class ChatPersistenceService {
 
         var msg = Message.builder()
                 .conversationId(conversationId)
+                .userKey(userKey)
                 .role(role)
                 .content(content)
                 .intent(intent)
@@ -51,22 +53,14 @@ public class ChatPersistenceService {
         messageRepository.save(msg);
     }
 
-    /**
-     * 分页查询对话历史。
-     * 使用数据库分页(Pageable)避免全量加载到内存导致 OOM,
-     * 并通过批量查询消息数解决 N+1 问题。
-     *
-     * @return Map 包含 items / page / page_size / total / total_pages
-     */
-    public Map<String, Object> getHistory(int page, int pageSize) {
-        // 防御性参数校验
+    public Map<String, Object> getHistory(String userKey, int page, int pageSize) {
         int safePage = Math.max(1, page);
         int safePageSize = Math.max(1, pageSize);
         Pageable pageable = PageRequest.of(safePage - 1, safePageSize);
 
-        Page<Conversation> convPage = conversationRepository.findAllByOrderByCreatedAtDesc(pageable);
+        Page<Conversation> convPage = conversationRepository
+                .findByUserKeyOrderByCreatedAtDesc(userKey, pageable);
 
-        // 批量查询消息数,解决 N+1 问题
         List<String> cids = convPage.getContent().stream()
                 .map(Conversation::getConversationId)
                 .collect(Collectors.toList());
@@ -123,4 +117,20 @@ public class ChatPersistenceService {
             messageRepository.save(msg);
         });
     }
-}
+
+    /** 获取当前用户最近 N 条消息,用于前端恢复对话 */
+    public List<Map<String, Object>> getRecentMessages(String userKey, int limit) {
+        List<Map<String, Object>> result = new ArrayList<>();
+        var page = messageRepository.findRecentByUserKey(userKey, PageRequest.of(0, limit));
+        var msgs = new ArrayList<>(page.getContent());
+        java.util.Collections.reverse(msgs);
+        for (var msg : msgs) {
+            result.add(Map.of(
+                    "role", msg.getRole(),
+                    "content", msg.getContent(),
+                    "intent", msg.getIntent() != null ? msg.getIntent() : ""
+            ));
+        }
+        return result;
+    }
+}

+ 58 - 14
static/index.html

@@ -1,4 +1,4 @@
-<!DOCTYPE html>
+<!DOCTYPE html>
 <html lang="zh-CN">
 <head>
   <meta charset="UTF-8">
@@ -98,8 +98,6 @@
     .pager button:disabled{opacity:.35;cursor:default}
     .pager .page-info{font-size:12px;color:#888;margin:0 4px;white-space:nowrap}
     .pager select{padding:6px 8px;appearance:auto}
-    .clear-chat-btn{position:absolute;right:12px;top:50%;transform:translateY(-50%);background:rgba(255,255,255,.15);border:1px solid rgba(255,255,255,.3);color:#fff;padding:4px 10px;border-radius:14px;font-size:11px;cursor:pointer;z-index:2;transition:all .15s}
-    .clear-chat-btn:hover{background:rgba(255,255,255,.25)}
     .login-modal-mask{position:fixed;inset:0;display:flex;align-items:center;justify-content:center;padding:24px;background:rgba(0,0,0,.5);z-index:100}
     .login-modal{width:min(320px,100%);overflow:hidden;border-radius:12px;background:#fff;box-shadow:0 12px 36px rgba(0,0,0,.2);text-align:center}
     .login-modal-title{padding:24px 24px 8px;color:var(--text);font-size:17px;font-weight:700}
@@ -115,7 +113,7 @@
 <body>
 <div class="header">
   <h1>AI 药典助手</h1>
-  <button class="clear-chat-btn" onclick="clearChat()" title="清除本地对话记录">🗑 清空对话</button>
+  
   <div class="sub">基于《中华人民共和国药典》& 通义千问</div>
 </div>
 <div class="nav">
@@ -198,6 +196,22 @@
     return h;
   }
 
+  // 自动获取 guest token(后端鉴权未就绪时静默降级)
+  async function ensureToken(){
+    if (TOKEN) return true;
+    try {
+      var r = await fetch(API_BASE + '/api/v1/auth/guest', {
+        method: 'POST', headers: {'Content-Type':'application/json'}
+      });
+      if (r.ok) {
+        var d = await r.json();
+        TOKEN = d.access_token || d.token || '';
+        if (TOKEN) { localStorage.setItem('pharma_token', TOKEN); return true; }
+      }
+    } catch(e) {}
+    return false;
+  }
+
   function classifyIntent(q){
     var t = q.trim();
     if (/怎么吃|吃多少|怎么用|孕妇|儿童用量|副作用|不良反应|禁忌|过敏|能不能|可以吗|安全吗|剂量/.test(t)) return 'usage_guide';
@@ -255,7 +269,7 @@
   // ============================================================
   // 本地对话持久化 (IndexedDB, 上限 256MB)
   // ============================================================
-  var chatDB=null,chatDBTotalSize=0,chatDBMaxSize=256*1024*1024;
+  var chatDB=null,chatDBTotalSize=0,chatDBMaxSize=128*1024*1024,chatDBMaxCount=500;
   function openChatDB(){
     return new Promise(function(resolve,reject){
       var req=indexedDB.open('PharmaChat',1);
@@ -274,9 +288,11 @@
     store.add(msg);
     chatDBTotalSize+=new Blob([JSON.stringify(msg)]).size;
     pruneIfNeeded();
+    pruneByCount();
   }
   function pruneIfNeeded(){
     if(chatDBTotalSize<=chatDBMaxSize)return;
+    // 淘汰最早的消息,直到低于 80% 上限
     var tx=chatDB.transaction('messages','readwrite');
     var store=tx.objectStore('messages');
     var req=store.openCursor();
@@ -289,6 +305,25 @@
       }
     };
   }
+  function pruneByCount(){
+    if(!chatDB)return;
+    var tx=chatDB.transaction('messages','readwrite');
+    var store=tx.objectStore('messages');
+    var countReq=store.count();
+    countReq.onsuccess=function(){
+      if(countReq.result<=chatDBMaxCount)return;
+      var excess=countReq.result-chatDBMaxCount;
+      var cursorReq=store.openCursor();
+      cursorReq.onsuccess=function(e){
+        var cursor=e.target.result;
+        if(cursor&&excess>0){
+          cursor.delete();
+          excess--;
+          cursor.continue();
+        }
+      };
+    };
+  }
   function loadChatFromDB(){
     return new Promise(function(resolve){
       if(!chatDB){resolve([]);return}
@@ -304,17 +339,20 @@
       req.onerror=function(){resolve([])};
     });
   }
-  window.clearChat=function(){
-    if(!chatDB)return;
-    if(!confirm('确定要清空所有本地对话记录吗?此操作不可撤销。'))return;
-    var tx=chatDB.transaction('messages','readwrite');
-    tx.objectStore('messages').clear();
-    chatDBTotalSize=0;
-    document.getElementById('chat').innerHTML='<div class="welcome"><div class="icon">📷💊🎬</div><p>输入药品名称或问题,AI 从药典知识库中检索回答<br><small style="color:#aaa">支持图片/视频上传识别 · 2025版药典</small></p></div>';
-  };
+  ;
   async function restoreChat(){
     var msgs=await loadChatFromDB();
+    // IndexedDB 为空时(微信 WebView 等场景),从后端恢复
+    if(!msgs.length){
+      try{
+        var r=await fetch(API_BASE+'/api/v1/chat/recent-messages?limit=50',{headers:authHeaders()});
+        if(r.ok){ var d=await r.json(); msgs=d.messages||[]; }
+      }catch(e){}
+    }
     if(!msgs.length)return;
+    renderMessages(msgs);
+  }
+  function renderMessages(msgs){
     var w=document.getElementById('chat').querySelector('.welcome');if(w)w.remove();
     for(var i=0;i<msgs.length;i++){
       var m=msgs[i],div=document.createElement('div');
@@ -481,10 +519,15 @@
     if(e.key==='Escape')closeLoginModal();
   });
 
-  window.send=function(){
+  window.send=async function(){
     var input=document.getElementById('userInput'),text=input.value.trim();
     var hasMedia=!!pendingMedia;
     if(!text&&!hasMedia)return;input.value='';if(STREAMING)return;
+    // 强鉴权:无 token 时自动获取,失败则弹登录框
+    if(!TOKEN){
+      var ok = await ensureToken();
+      if(!ok){ showLoginModal(); return; }
+    }
     doSend(text,hasMedia?pendingMedia:null);
   };
   function doSend(text,media){
@@ -650,6 +693,7 @@
   };
   // 初始化:打开本地数据库 → 恢复对话 → 预加载药品列表
       (async function init(){
+        await ensureToken();
         await openChatDB();
         await restoreChat();
         loadDrugList();