package com.superbiz.agent.tool; import org.springframework.stereotype.Component; import java.util.Collections; import java.util.List; import java.util.Map; import java.util.Set; import java.util.concurrent.ConcurrentHashMap; import java.util.stream.Collectors; /** * session 级已召回文档追踪器 * 支持文档级去重 + 域级行动记忆 * * 数据结构:sessionId → { domain → Set } * - 域级:控制"不要重复查同域",提供行动记忆给 LLM * - 文档级:控制"不要重复召回同文档"(替代原有单层结构) */ @Component public class RetrievedDocTracker { // key: sessionId, value: { domain → Set } private final ConcurrentHashMap>> sessionRetrievals = new ConcurrentHashMap<>(); /** * 记录一次检索(域级 + 文档级) */ public void markRetrieved(String sessionId, String domain, String filePath) { if (sessionId == null || filePath == null) return; sessionRetrievals.computeIfAbsent(sessionId, k -> new ConcurrentHashMap<>()) .computeIfAbsent(domain != null ? domain : "_unknown", d -> Collections.newSetFromMap(new ConcurrentHashMap<>())) .add(filePath); } /** * 文档级去重:检查 filePath 是否已在本会话中检索过 */ public boolean isDocRetrieved(String sessionId, String filePath) { if (sessionId == null || filePath == null) return false; Map> domains = sessionRetrievals.get(sessionId); if (domains == null) return false; return domains.values().stream().anyMatch(docs -> docs.contains(filePath)); } /** * 域级检查:检查 domain 是否已在本会话中检索过 */ public boolean isDomainRetrieved(String sessionId, String domain) { if (sessionId == null || domain == null) return false; Map> domains = sessionRetrievals.get(sessionId); return domains != null && domains.containsKey(domain); } /** * 获取本次会话已检索的域列表(行动记忆,返回给 LLM) */ public List getRetrievedDomains(String sessionId) { if (sessionId == null) return List.of(); Map> domains = sessionRetrievals.get(sessionId); if (domains == null) return List.of(); return List.copyOf(domains.keySet()); } /** * 向后兼容:文档级去重(委托给 isDocRetrieved) */ public boolean isAlreadyRetrieved(String sessionId, String filePath) { return isDocRetrieved(sessionId, filePath); } /** * 向后兼容:旧版 markRetrieved(domain 设为 null,归入 _unknown) */ public void markRetrieved(String sessionId, String filePath) { markRetrieved(sessionId, null, filePath); } /** * 清理会话 */ public void clearSession(String sessionId) { if (sessionId == null) return; sessionRetrievals.remove(sessionId); } }