diff --git a/src/main/java/com/knowledge/base/application/filter/AuthFilter.java b/src/main/java/com/knowledge/base/application/filter/AuthFilter.java index 505a293..0a42d33 100644 --- a/src/main/java/com/knowledge/base/application/filter/AuthFilter.java +++ b/src/main/java/com/knowledge/base/application/filter/AuthFilter.java @@ -33,7 +33,9 @@ public class AuthFilter implements Filter { String path = request.getRequestURI(); // 需要鉴权的路径前缀 - if (path.startsWith("/api/v1/doc") || path.startsWith("/api/v1/llm")) { + if (path.startsWith("/api/v1/doc") +// || path.startsWith("/api/v1/llm") + ) { String token = request.getHeader("Authorization"); if (StrUtil.isBlank(token)) { diff --git a/src/main/java/com/knowledge/base/application/service/LLMAppService.java b/src/main/java/com/knowledge/base/application/service/LLMAppService.java index 89e9130..55dd490 100644 --- a/src/main/java/com/knowledge/base/application/service/LLMAppService.java +++ b/src/main/java/com/knowledge/base/application/service/LLMAppService.java @@ -5,7 +5,22 @@ import org.springframework.web.servlet.mvc.method.annotation.SseEmitter; import java.util.Map; public interface LLMAppService { + /** + * 获取llmtoken + * @param password + * @return + * @throws Exception + */ String getToken(String password) throws Exception; + + /** + * 问llm问题 + * @param llmToken + * @param question + * @param params + * @return + * @throws Exception + */ SseEmitter ask(String llmToken, String question, Map params) throws Exception; } diff --git a/src/main/java/com/knowledge/base/application/service/LLMAppServiceImpl.java b/src/main/java/com/knowledge/base/application/service/LLMAppServiceImpl.java index e1c6a02..4807359 100644 --- a/src/main/java/com/knowledge/base/application/service/LLMAppServiceImpl.java +++ b/src/main/java/com/knowledge/base/application/service/LLMAppServiceImpl.java @@ -1,7 +1,11 @@ package com.knowledge.base.application.service; +import com.knowledge.base.infrastructure.config.ThreadPoolConfig; import com.knowledge.base.infrastructure.south.llm.LLMServiceFactory; +import com.knowledge.base.infrastructure.util.RateLimiterManager; +import com.knowledge.base.infrastructure.util.ThreadPoolUtil; import com.knowledge.base.infrastructure.util.http.FilteredSseOutputAdapter; +import com.knowledge.base.infrastructure.util.http.WriterAdapter; import lombok.RequiredArgsConstructor; import lombok.extern.slf4j.Slf4j; import org.springframework.stereotype.Service; @@ -17,6 +21,8 @@ public class LLMAppServiceImpl implements LLMAppService { private final LLMServiceFactory llmServiceFactory; + private final RateLimiterManager rateLimiterManager; + @Override public String getToken(String password) throws Exception { return llmServiceFactory.current().fetchToken(password); @@ -24,12 +30,14 @@ public class LLMAppServiceImpl implements LLMAppService { @Override public SseEmitter ask(String llmToken, String question, Map params) throws Exception { - SseEmitter emitter = new SseEmitter(0L); // 不设超时 + SseEmitter emitter = new SseEmitter(300 * 1000L); // 超时时间设为5分钟 - new Thread(() -> { + ThreadPoolUtil.execute(() -> { try { - FilteredSseOutputAdapter adapter = new FilteredSseOutputAdapter(emitter); + rateLimiterManager.getRateLimiter(RateLimiterManager.RATE_LIMIT_SCENE_LLM_ASK).acquire(); + WriterAdapter adapter = new FilteredSseOutputAdapter(emitter); llmServiceFactory.current().streamAnswer(llmToken, question, params, adapter); + emitter.complete(); } catch (Exception e) { log.error("LLM调用异常", e); try { @@ -39,7 +47,7 @@ public class LLMAppServiceImpl implements LLMAppService { log.warn("SSE发送错误信息失败", ioException); } } - }).start(); + }, ThreadPoolConfig.SSE_POOL); return emitter; } diff --git a/src/main/java/com/knowledge/base/domain/doc/service/impl/importer/AbstractBaseFileImporter.java b/src/main/java/com/knowledge/base/domain/doc/service/impl/importer/AbstractBaseFileImporter.java index 6cc139c..4f9d82c 100644 --- a/src/main/java/com/knowledge/base/domain/doc/service/impl/importer/AbstractBaseFileImporter.java +++ b/src/main/java/com/knowledge/base/domain/doc/service/impl/importer/AbstractBaseFileImporter.java @@ -64,7 +64,7 @@ public abstract class AbstractBaseFileImporter implements DocumentImporter { return getFileSuffixes().stream().anyMatch(fileName::endsWith); }) .forEach(path -> { - rateLimiterManager.getRateLimiter().acquire(); + rateLimiterManager.getRateLimiter(RateLimiterManager.RATE_LIMIT_SCENE_IMPORT).acquire(); ThreadPoolUtil.execute(() -> insertOrUpdateOneFileIntoES(path, excludeBase, Maps.newHashMap()), ThreadPoolConfig.IMPORT_DOC_POOL); }); diff --git a/src/main/java/com/knowledge/base/infrastructure/config/DynamicConfig.java b/src/main/java/com/knowledge/base/infrastructure/config/DynamicConfig.java index 83954e7..317e17f 100644 --- a/src/main/java/com/knowledge/base/infrastructure/config/DynamicConfig.java +++ b/src/main/java/com/knowledge/base/infrastructure/config/DynamicConfig.java @@ -34,4 +34,7 @@ public class DynamicConfig { @Value("${file.import.rate.limit: 10}") private String fileImportRateLimit; + + @Value("${llm.sse.rate.limit: 3}") + private String llmSseRateLimit; } diff --git a/src/main/java/com/knowledge/base/infrastructure/config/ThreadPoolConfig.java b/src/main/java/com/knowledge/base/infrastructure/config/ThreadPoolConfig.java index 8bf0a10..957e1c3 100644 --- a/src/main/java/com/knowledge/base/infrastructure/config/ThreadPoolConfig.java +++ b/src/main/java/com/knowledge/base/infrastructure/config/ThreadPoolConfig.java @@ -18,4 +18,9 @@ public class ThreadPoolConfig { ThreadFactoryBuilder.create().setNamePrefix("Import-Doc-pool-").build(), new ThreadPoolExecutor.CallerRunsPolicy() ); + + public static final ExecutorService SSE_POOL = new ThreadPoolExecutor( + 10, 50, 60L, TimeUnit.SECONDS, new LinkedBlockingQueue<>(1000), // 可根据实际调优 + new ThreadPoolExecutor.AbortPolicy() + ); } diff --git a/src/main/java/com/knowledge/base/infrastructure/north/controller/LLMController.java b/src/main/java/com/knowledge/base/infrastructure/north/controller/LLMController.java index bfc1889..960244b 100644 --- a/src/main/java/com/knowledge/base/infrastructure/north/controller/LLMController.java +++ b/src/main/java/com/knowledge/base/infrastructure/north/controller/LLMController.java @@ -2,15 +2,16 @@ package com.knowledge.base.infrastructure.north.controller; import com.knowledge.base.application.service.LLMAppService; import lombok.RequiredArgsConstructor; +import lombok.extern.slf4j.Slf4j; import org.springframework.http.MediaType; import org.springframework.web.bind.annotation.*; import org.springframework.web.servlet.mvc.method.annotation.SseEmitter; - import java.util.Map; @RestController @RequestMapping("/api/v1/llm") @RequiredArgsConstructor +@Slf4j public class LLMController { private final LLMAppService llmAppService; diff --git a/src/main/java/com/knowledge/base/infrastructure/south/llm/AnythingLLMServiceImpl.java b/src/main/java/com/knowledge/base/infrastructure/south/llm/AnythingLLMServiceImpl.java index 055f9fc..b790a0f 100644 --- a/src/main/java/com/knowledge/base/infrastructure/south/llm/AnythingLLMServiceImpl.java +++ b/src/main/java/com/knowledge/base/infrastructure/south/llm/AnythingLLMServiceImpl.java @@ -15,7 +15,6 @@ import org.springframework.stereotype.Service; import java.io.BufferedReader; import java.io.InputStreamReader; -import java.io.OutputStream; import java.nio.charset.StandardCharsets; import java.util.HashMap; import java.util.List; diff --git a/src/main/java/com/knowledge/base/infrastructure/util/RateLimiterManager.java b/src/main/java/com/knowledge/base/infrastructure/util/RateLimiterManager.java index 5ebc672..0a2ffc4 100644 --- a/src/main/java/com/knowledge/base/infrastructure/util/RateLimiterManager.java +++ b/src/main/java/com/knowledge/base/infrastructure/util/RateLimiterManager.java @@ -5,28 +5,72 @@ import com.knowledge.base.infrastructure.config.DynamicConfig; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.stereotype.Component; +import java.util.Map; +import java.util.concurrent.ConcurrentHashMap; + @Component public class RateLimiterManager { @Autowired private DynamicConfig dynamicConfig; - private volatile RateLimiter rateLimiter; - private volatile double lastRate = -1; + public static final String RATE_LIMIT_SCENE_IMPORT = "fileImport"; + public static final String RATE_LIMIT_SCENE_LLM_ASK = "llmAsk"; - public RateLimiter getRateLimiter() { - double currentRate = Double.valueOf(dynamicConfig.getFileImportRateLimit()); + // 每个场景一个独立 RateLimiter + private final Map limiterMap = new ConcurrentHashMap<>(); - // 如果速率发生变化,则更新限速器 - if (rateLimiter == null || currentRate != lastRate) { - synchronized (this) { - if (rateLimiter == null || currentRate != lastRate) { - rateLimiter = RateLimiter.create(currentRate); - lastRate = currentRate; - } + /** + * 获取指定场景的限流器,支持动态配置速率 + * @param scene 场景名称,如 "fileImport", "llmSse" + * @return 对应 RateLimiter + */ + public RateLimiter getRateLimiter(String scene) { + String rateKey = getRateKeyForScene(scene); + double currentRate = getConfigRate(rateKey); + + return limiterMap.compute(scene, (k, holder) -> { + if (holder == null || holder.rate != currentRate) { + return new RateLimiterHolder(RateLimiter.create(currentRate), currentRate); } - } + return holder; + }).rateLimiter; + } - return rateLimiter; + private String getRateKeyForScene(String scene) { + switch (scene) { + case RATE_LIMIT_SCENE_IMPORT: + return "fileImportRateLimit"; + case RATE_LIMIT_SCENE_LLM_ASK: + return "llmSseRateLimit"; + default: + return "defaultRateLimit"; + } + } + + private double getConfigRate(String rateKey) { + try { + switch (rateKey) { + case "fileImportRateLimit": + return Double.parseDouble(dynamicConfig.getFileImportRateLimit()); + case "llmSseRateLimit": + return Double.parseDouble(dynamicConfig.getLlmSseRateLimit()); + default: + return 1.0; + } + } catch (Exception e) { + return 1.0; + } + } + + + private static class RateLimiterHolder { + final RateLimiter rateLimiter; + // nacos配置变化,需更改 + final double rate; + RateLimiterHolder(RateLimiter rl, double rate) { + this.rateLimiter = rl; + this.rate = rate; + } } } diff --git a/src/main/java/com/knowledge/base/infrastructure/util/http/FilteredSseOutputAdapter.java b/src/main/java/com/knowledge/base/infrastructure/util/http/FilteredSseOutputAdapter.java index 59726d9..7bbcd6d 100644 --- a/src/main/java/com/knowledge/base/infrastructure/util/http/FilteredSseOutputAdapter.java +++ b/src/main/java/com/knowledge/base/infrastructure/util/http/FilteredSseOutputAdapter.java @@ -1,8 +1,10 @@ package com.knowledge.base.infrastructure.util.http; +import cn.hutool.core.map.MapUtil; import cn.hutool.core.util.StrUtil; import com.fasterxml.jackson.databind.ObjectMapper; import lombok.extern.slf4j.Slf4j; +import org.springframework.http.MediaType; import org.springframework.web.servlet.mvc.method.annotation.SseEmitter; import java.io.IOException; @@ -35,33 +37,29 @@ public class FilteredSseOutputAdapter implements WriterAdapter { try { Map original = mapper.readValue(line, Map.class); Map filtered = new HashMap<>(); - - if (original.containsKey("textResponse")) { - filtered.put("textResponse", original.get("textResponse")); - lastChunk.put("textResponse", original.get("textResponse")); - } - - if (original.containsKey("sources")) { - lastChunk.put("sources", original.get("sources")); - } - - if (original.containsKey("close")) { - lastChunk.put("close", original.get("close")); + if(MapUtil.isEmpty(original)) { + log.warn("original map is empty"); + return; } if (Boolean.TRUE.equals(original.get("close"))) { - filtered.putAll(lastChunk); - lastChunk.clear(); - + // 最后一条特殊处理 + filtered.put("textResponse", original.get("textResponse")); + filtered.put("sources", original.get("sources")); + filtered.put("close", Boolean.TRUE); String payload = mapper.writeValueAsString(filtered); - emitter.send(SseEmitter.event().data(payload)); + log.debug("发送SSE段: {}", payload); + emitter.send(SseEmitter.event().data(payload, MediaType.TEXT_EVENT_STREAM)); closed = true; - emitter.complete(); - } else if (!filtered.isEmpty()) { - String payload = mapper.writeValueAsString(filtered); - emitter.send(SseEmitter.event().data(payload)); - } + } else { + // 其它统一只返回textResponse + filtered.put("textResponse", original.get("textResponse")); + filtered.put("close", Boolean.FALSE); + String payload = mapper.writeValueAsString(filtered); + log.debug("发送SSE段: {}", payload); + emitter.send(SseEmitter.event().data(payload, MediaType.TEXT_EVENT_STREAM)); + } } catch (IllegalStateException e) { log.warn("SSE连接已关闭,忽略发送: {}", e.getMessage()); closed = true;