yudao-module-ai/src/main/java/cn/iocoder/yudao/module/ai/framework/ai/core/model/AiModelFactoryImpl.java
@@ -1,281 +1,59 @@ package cn.iocoder.yudao.module.ai.framework.ai.core.model; import cn.hutool.core.io.FileUtil; import cn.hutool.core.lang.Assert; import cn.hutool.core.lang.Singleton; import cn.hutool.core.lang.func.Func0; import cn.hutool.core.util.ArrayUtil; import cn.hutool.core.util.RuntimeUtil; import cn.hutool.core.util.StrUtil; import cn.hutool.extra.spring.SpringUtil; import cn.iocoder.yudao.module.ai.enums.model.AiPlatformEnum; import cn.iocoder.yudao.module.ai.framework.ai.config.AiAutoConfiguration; import cn.iocoder.yudao.module.ai.framework.ai.config.YudaoAiProperties; import cn.iocoder.yudao.module.ai.framework.ai.core.model.baichuan.BaiChuanChatModel; import cn.iocoder.yudao.module.ai.framework.ai.core.model.doubao.DouBaoChatModel; import cn.iocoder.yudao.module.ai.framework.ai.core.model.hunyuan.HunYuanChatModel; import cn.iocoder.yudao.module.ai.framework.ai.core.model.midjourney.api.MidjourneyApi; import cn.iocoder.yudao.module.ai.framework.ai.core.model.minimax.MiniMaxChatModel; import cn.iocoder.yudao.module.ai.framework.ai.core.model.moonshot.MoonshotChatModel; import cn.iocoder.yudao.module.ai.framework.ai.core.model.siliconflow.SiliconFlowApiConstants; import cn.iocoder.yudao.module.ai.framework.ai.core.model.siliconflow.SiliconFlowChatModel; import cn.iocoder.yudao.module.ai.framework.ai.core.model.siliconflow.SiliconFlowImageApi; import cn.iocoder.yudao.module.ai.framework.ai.core.model.siliconflow.SiliconFlowImageModel; import cn.iocoder.yudao.module.ai.framework.ai.core.model.stepfun.StepFunChatModel; import cn.iocoder.yudao.module.ai.framework.ai.core.model.suno.api.SunoApi; import cn.iocoder.yudao.module.ai.framework.ai.core.model.xinghuo.XingHuoChatModel; import cn.iocoder.yudao.module.ai.framework.ai.core.model.yiyan.YiYanChatModel; import cn.iocoder.yudao.module.ai.framework.ai.core.model.zhipu.ZhiPuChatModel; import cn.iocoder.yudao.module.ai.util.AiUtils; import com.alibaba.cloud.ai.autoconfigure.dashscope.DashScopeChatAutoConfiguration; import com.alibaba.cloud.ai.autoconfigure.dashscope.DashScopeEmbeddingAutoConfiguration; import com.alibaba.cloud.ai.autoconfigure.dashscope.DashScopeImageAutoConfiguration; import com.alibaba.cloud.ai.dashscope.chat.DashScopeChatModel; import com.alibaba.cloud.ai.dashscope.embedding.text.DashScopeEmbeddingModel; import com.alibaba.cloud.ai.dashscope.image.DashScopeImageModel; import com.google.genai.Client; import com.google.genai.types.HttpOptions; import io.micrometer.observation.ObservationRegistry; import io.milvus.client.MilvusServiceClient; import io.qdrant.client.QdrantClient; import io.qdrant.client.QdrantGrpcClient; import lombok.SneakyThrows; import org.springframework.ai.anthropic.AnthropicChatModel; import org.springframework.ai.anthropic.AnthropicChatOptions; import lombok.extern.slf4j.Slf4j; import org.springframework.ai.chat.model.ChatModel; import org.springframework.ai.deepseek.DeepSeekChatModel; import org.springframework.ai.deepseek.DeepSeekChatOptions; import org.springframework.ai.deepseek.api.DeepSeekApi; import org.springframework.ai.document.MetadataMode; import org.springframework.ai.embedding.BatchingStrategy; import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.embedding.observation.EmbeddingModelObservationConvention; import org.springframework.ai.google.genai.GoogleGenAiChatModel; import org.springframework.ai.google.genai.GoogleGenAiChatOptions; import org.springframework.ai.image.ImageModel; import org.springframework.ai.model.anthropic.autoconfigure.AnthropicChatAutoConfiguration; import org.springframework.ai.model.deepseek.autoconfigure.DeepSeekChatAutoConfiguration; import org.springframework.ai.model.google.genai.autoconfigure.chat.GoogleGenAiChatAutoConfiguration; import org.springframework.ai.model.ollama.autoconfigure.OllamaChatAutoConfiguration; import org.springframework.ai.model.openai.autoconfigure.OpenAiChatAutoConfiguration; import org.springframework.ai.model.openai.autoconfigure.OpenAiEmbeddingAutoConfiguration; import org.springframework.ai.model.openai.autoconfigure.OpenAiImageAutoConfiguration; import org.springframework.ai.model.stabilityai.autoconfigure.StabilityAiImageAutoConfiguration; import org.springframework.ai.model.tool.ToolCallingManager; import org.springframework.ai.ollama.OllamaChatModel; import org.springframework.ai.ollama.OllamaEmbeddingModel; import org.springframework.ai.ollama.api.OllamaApi; import org.springframework.ai.ollama.api.OllamaEmbeddingOptions; import org.springframework.ai.openai.*; import org.springframework.ai.retry.RetryUtils; import org.springframework.ai.stabilityai.StabilityAiImageModel; import org.springframework.ai.stabilityai.api.StabilityAiApi; import org.springframework.ai.vectorstore.SimpleVectorStore; import org.springframework.ai.embedding.TokenCountBatchingStrategy; import org.springframework.ai.openai.OpenAiChatModel; import org.springframework.ai.vectorstore.VectorStore; import org.springframework.ai.vectorstore.milvus.MilvusVectorStore; import org.springframework.ai.vectorstore.milvus.autoconfigure.MilvusServiceClientConnectionDetails; import org.springframework.ai.vectorstore.milvus.autoconfigure.MilvusServiceClientProperties; import org.springframework.ai.vectorstore.milvus.autoconfigure.MilvusVectorStoreAutoConfiguration; import org.springframework.ai.vectorstore.milvus.autoconfigure.MilvusVectorStoreProperties; import org.springframework.ai.vectorstore.observation.DefaultVectorStoreObservationConvention; import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention; import org.springframework.ai.vectorstore.qdrant.QdrantVectorStore; import org.springframework.ai.vectorstore.qdrant.autoconfigure.QdrantVectorStoreAutoConfiguration; import org.springframework.ai.vectorstore.qdrant.autoconfigure.QdrantVectorStoreProperties; import org.springframework.ai.vectorstore.redis.RedisVectorStore; import org.springframework.ai.vectorstore.redis.autoconfigure.RedisVectorStoreAutoConfiguration; import org.springframework.ai.vectorstore.redis.autoconfigure.RedisVectorStoreProperties; import org.springframework.beans.BeansException; import org.springframework.beans.factory.ObjectProvider; import org.springframework.boot.data.redis.autoconfigure.DataRedisProperties; import redis.clients.jedis.DefaultJedisClientConfig; import redis.clients.jedis.HostAndPort; import redis.clients.jedis.JedisClientConfig; import redis.clients.jedis.RedisClient; import java.io.File; import java.time.Duration; import java.util.Collections; import java.util.Map; import java.util.Timer; import java.util.TimerTask; import static cn.iocoder.yudao.framework.common.util.collection.CollectionUtils.convertList; /** * AI Model 模型工厂的实现类 * AI Model 模型工厂实现类 * * @author 芋道源码 * 使用 OpenAI 兼容接口对接通义千问 + Milvus 向量存储 */ @Slf4j public class AiModelFactoryImpl implements AiModelFactory { @Override public ChatModel getOrCreateChatModel(AiPlatformEnum platform, String rawApiKey, String rawUrl) { final String apiKey = resolveSpringPlaceholders(rawApiKey); final String url = resolveSpringPlaceholders(rawUrl); String cacheKey = buildClientCacheKey(ChatModel.class, platform, apiKey, url); public ChatModel getOrCreateChatModel(AiPlatformEnum platform, String apiKey, String url, String model) { String cacheKey = buildCacheKey(ChatModel.class, platform, apiKey, url, model); return Singleton.get(cacheKey, (Func0<ChatModel>) () -> { // noinspection EnhancedSwitchMigration switch (platform) { case TONG_YI: return buildTongYiChatModel(apiKey); case YI_YAN: return buildYiYanChatModel(apiKey); case DEEP_SEEK: return buildDeepSeekChatModel(apiKey); case DOU_BAO: return buildDouBaoChatModel(apiKey); case HUN_YUAN: return buildHunYuanChatModel(apiKey, url); case SILICON_FLOW: return buildSiliconFlowChatModel(apiKey); case ZHI_PU: return buildZhiPuChatModel(apiKey, url); case MINI_MAX: return buildMiniMaxChatModel(apiKey, url); case MOONSHOT: return buildMoonshotChatModel(apiKey, url); case STEP_FUN: return buildStepFunChatModel(apiKey, url); case XING_HUO: return buildXingHuoChatModel(apiKey); case BAI_CHUAN: return buildBaiChuanChatModel(apiKey); case OPENAI: return buildOpenAiChatModel(apiKey, url); case AZURE_OPENAI: return buildAzureOpenAiChatModel(apiKey, url); case ANTHROPIC: return buildAnthropicChatModel(apiKey, url); case GEMINI: return buildGeminiChatModel(apiKey, url); case OLLAMA: return buildOllamaChatModel(url); case GROK: return buildGrokChatModel(apiKey, url); default: throw new IllegalArgumentException(StrUtil.format("未知平台({})", platform)); if (platform == AiPlatformEnum.TONG_YI) { return AiAutoConfiguration.buildTongYiChatModel(apiKey, model); } throw new IllegalArgumentException(StrUtil.format("不支持的平台({})", platform)); }); } @Override public ChatModel getDefaultChatModel(AiPlatformEnum platform) { // noinspection EnhancedSwitchMigration switch (platform) { case TONG_YI: return SpringUtil.getBean(DashScopeChatModel.class); case YI_YAN: return SpringUtil.getBean(YiYanChatModel.class); case DEEP_SEEK: return SpringUtil.getBean(DeepSeekChatModel.class); case DOU_BAO: return SpringUtil.getBean(DouBaoChatModel.class); case HUN_YUAN: return SpringUtil.getBean(HunYuanChatModel.class); case SILICON_FLOW: return SpringUtil.getBean(SiliconFlowChatModel.class); case ZHI_PU: return SpringUtil.getBean(ZhiPuChatModel.class); case MINI_MAX: return SpringUtil.getBean(MiniMaxChatModel.class); case MOONSHOT: return SpringUtil.getBean(MoonshotChatModel.class); case STEP_FUN: return SpringUtil.getBean(StepFunChatModel.class); case XING_HUO: return SpringUtil.getBean(XingHuoChatModel.class); case BAI_CHUAN: return SpringUtil.getBean(BaiChuanChatModel.class); case OPENAI: return SpringUtil.getBean(OpenAiChatModel.class); case ANTHROPIC: return SpringUtil.getBean(AnthropicChatModel.class); case GEMINI: return SpringUtil.getBean(GoogleGenAiChatModel.class); case OLLAMA: return SpringUtil.getBean(OllamaChatModel.class); default: throw new IllegalArgumentException(StrUtil.format("未知平台({})", platform)); if (platform == AiPlatformEnum.TONG_YI) { return SpringUtil.getBean(OpenAiChatModel.class); } throw new IllegalArgumentException(StrUtil.format("不支持的平台({})", platform)); } @Override public ImageModel getDefaultImageModel(AiPlatformEnum platform) { // noinspection EnhancedSwitchMigration switch (platform) { case TONG_YI: return SpringUtil.getBean(DashScopeImageModel.class); case SILICON_FLOW: return SpringUtil.getBean(SiliconFlowImageModel.class); case OPENAI: return SpringUtil.getBean(OpenAiImageModel.class); case STABLE_DIFFUSION: return SpringUtil.getBean(StabilityAiImageModel.class); default: throw new IllegalArgumentException(StrUtil.format("未知平台({})", platform)); } } @Override public ImageModel getOrCreateImageModel(AiPlatformEnum platform, String rawApiKey, String rawUrl) { String apiKey = resolveSpringPlaceholders(rawApiKey); String url = resolveSpringPlaceholders(rawUrl); // noinspection EnhancedSwitchMigration switch (platform) { case TONG_YI: return buildTongYiImagesModel(apiKey); case OPENAI: return buildOpenAiImageModel(apiKey, url); case SILICON_FLOW: return buildSiliconFlowImageModel(apiKey, url); case STABLE_DIFFUSION: return buildStabilityAiImageModel(apiKey, url); default: throw new IllegalArgumentException(StrUtil.format("未知平台({})", platform)); } } @Override public MidjourneyApi getOrCreateMidjourneyApi(String rawApiKey, String rawUrl) { final String apiKey = resolveSpringPlaceholders(rawApiKey); final String url = resolveSpringPlaceholders(rawUrl); String cacheKey = buildClientCacheKey(MidjourneyApi.class, AiPlatformEnum.MIDJOURNEY.getPlatform(), apiKey, url); return Singleton.get(cacheKey, (Func0<MidjourneyApi>) () -> { YudaoAiProperties.Midjourney properties = SpringUtil.getBean(YudaoAiProperties.class) .getMidjourney(); return new MidjourneyApi(url, apiKey, properties.getNotifyUrl()); }); } @Override public SunoApi getOrCreateSunoApi(String rawApiKey, String rawUrl) { final String apiKey = resolveSpringPlaceholders(rawApiKey); final String url = resolveSpringPlaceholders(rawUrl); String cacheKey = buildClientCacheKey(SunoApi.class, AiPlatformEnum.SUNO.getPlatform(), apiKey, url); return Singleton.get(cacheKey, (Func0<SunoApi>) () -> new SunoApi(url)); } @Override @SuppressWarnings("EnhancedSwitchMigration") public EmbeddingModel getOrCreateEmbeddingModel(AiPlatformEnum platform, String rawApiKey, String rawUrl, String model) { final String apiKey = resolveSpringPlaceholders(rawApiKey); final String url = resolveSpringPlaceholders(rawUrl); String cacheKey = buildClientCacheKey(EmbeddingModel.class, platform, apiKey, url, model); public EmbeddingModel getOrCreateEmbeddingModel(AiPlatformEnum platform, String apiKey, String url, String model) { String cacheKey = buildCacheKey(EmbeddingModel.class, platform, apiKey, url, model); return Singleton.get(cacheKey, (Func0<EmbeddingModel>) () -> { switch (platform) { case TONG_YI: return buildTongYiEmbeddingModel(apiKey, model); case OPENAI: return buildOpenAiEmbeddingModel(apiKey, url, model); case AZURE_OPENAI: return buildAzureOpenAiEmbeddingModel(apiKey, url, model); case OLLAMA: return buildOllamaEmbeddingModel(url, model); default: throw new IllegalArgumentException(StrUtil.format("未知平台({})", platform)); if (platform == AiPlatformEnum.TONG_YI) { return AiAutoConfiguration.buildTongYiEmbeddingModel(apiKey, model); } throw new IllegalArgumentException(StrUtil.format("不支持的平台({})", platform)); }); } @@ -283,504 +61,44 @@ public VectorStore getOrCreateVectorStore(Class<? extends VectorStore> type, EmbeddingModel embeddingModel, Map<String, Class<?>> metadataFields) { String cacheKey = buildClientCacheKey(VectorStore.class, embeddingModel, type); // metadataFields 参与缓存 key,确保不同知识库使用不同配置时不会复用 String cacheKey = buildCacheKey(VectorStore.class, embeddingModel, type, metadataFields.hashCode()); return Singleton.get(cacheKey, (Func0<VectorStore>) () -> { if (type == SimpleVectorStore.class) { return buildSimpleVectorStore(embeddingModel); } if (type == QdrantVectorStore.class) { return buildQdrantVectorStore(embeddingModel); } if (type == RedisVectorStore.class) { return buildRedisVectorStore(embeddingModel, metadataFields); } if (type == MilvusVectorStore.class) { return buildMilvusVectorStore(embeddingModel); } throw new IllegalArgumentException(StrUtil.format("未知类型({})", type)); throw new IllegalArgumentException(StrUtil.format("不支持的向量存储类型({})", type)); }); } private static String buildClientCacheKey(Class<?> clazz, Object... params) { if (ArrayUtil.isEmpty(params)) { return clazz.getName(); } return StrUtil.format("{}#{}", clazz.getName(), ArrayUtil.join(params, "_")); } private static String resolveSpringPlaceholders(String value) { // yml 配置的占位符由 Spring 自动解析;DB 里保存的 ${xxx} 需要在这里手动解析。 return AiUtils.resolveSpringPlaceholders(value); } // ========== 各种创建 spring-ai 客户端的方法 ========== /** * 可参考 {@link DashScopeChatAutoConfiguration} 的 dashscopeChatModel 方法 */ private static DashScopeChatModel buildTongYiChatModel(String key) { return AiAutoConfiguration.buildTongYiChatModel(key); } /** * 可参考 {@link DashScopeImageAutoConfiguration} 的 dashScopeImageModel 方法 */ private static DashScopeImageModel buildTongYiImagesModel(String key) { return AiAutoConfiguration.buildTongYiImagesModel(key); } private ChatModel buildYiYanChatModel(String apiKey) { YudaoAiProperties.YiYan properties = new YudaoAiProperties.YiYan() .setApiKey(apiKey); return new AiAutoConfiguration().buildYiYanChatClient(properties); } /** * 可参考 {@link DeepSeekChatAutoConfiguration} 的 deepSeekChatModel 方法 */ private static DeepSeekChatModel buildDeepSeekChatModel(String apiKey) { DeepSeekApi deepSeekApi = DeepSeekApi.builder().apiKey(apiKey).build(); DeepSeekChatOptions options = DeepSeekChatOptions.builder().model(DeepSeekApi.DEFAULT_CHAT_MODEL) .temperature(0.7).build(); return DeepSeekChatModel.builder() .deepSeekApi(deepSeekApi) .options(options) .build(); } /** * 可参考 {@link AiAutoConfiguration#douBaoChatClient(YudaoAiProperties)} */ private ChatModel buildDouBaoChatModel(String apiKey) { YudaoAiProperties.DouBao properties = new YudaoAiProperties.DouBao() .setApiKey(apiKey); return new AiAutoConfiguration().buildDouBaoChatClient(properties); } /** * 可参考 {@link AiAutoConfiguration#hunYuanChatClient(YudaoAiProperties)} */ private ChatModel buildHunYuanChatModel(String apiKey, String url) { YudaoAiProperties.HunYuan properties = new YudaoAiProperties.HunYuan() .setBaseUrl(url).setApiKey(apiKey); return new AiAutoConfiguration().buildHunYuanChatClient(properties); } /** * 可参考 {@link AiAutoConfiguration#siliconFlowChatClient(YudaoAiProperties)} */ private ChatModel buildSiliconFlowChatModel(String apiKey) { YudaoAiProperties.SiliconFlow properties = new YudaoAiProperties.SiliconFlow() .setApiKey(apiKey); return new AiAutoConfiguration().buildSiliconFlowChatClient(properties); } /** * 可参考 {@link AiAutoConfiguration#zhiPuChatClient(YudaoAiProperties)} */ private ZhiPuChatModel buildZhiPuChatModel(String apiKey, String url) { YudaoAiProperties.ZhiPu properties = new YudaoAiProperties.ZhiPu() .setBaseUrl(url).setApiKey(apiKey); return new AiAutoConfiguration().buildZhiPuChatClient(properties); } /** * 可参考 {@link AiAutoConfiguration#miniMaxChatClient(YudaoAiProperties)} */ private MiniMaxChatModel buildMiniMaxChatModel(String apiKey, String url) { YudaoAiProperties.MiniMax properties = new YudaoAiProperties.MiniMax() .setBaseUrl(url).setApiKey(apiKey); return new AiAutoConfiguration().buildMiniMaxChatClient(properties); } /** * 可参考 {@link AiAutoConfiguration#moonshotChatClient(YudaoAiProperties)} */ private MoonshotChatModel buildMoonshotChatModel(String apiKey, String url) { YudaoAiProperties.Moonshot properties = new YudaoAiProperties.Moonshot() .setBaseUrl(url).setApiKey(apiKey); return new AiAutoConfiguration().buildMoonshotChatClient(properties); } /** * 可参考 {@link AiAutoConfiguration#stepFunChatClient(YudaoAiProperties)} */ private StepFunChatModel buildStepFunChatModel(String apiKey, String url) { YudaoAiProperties.StepFun properties = new YudaoAiProperties.StepFun() .setBaseUrl(url).setApiKey(apiKey); return new AiAutoConfiguration().buildStepFunChatClient(properties); } /** * 可参考 {@link AiAutoConfiguration#xingHuoChatClient(YudaoAiProperties)} */ private static XingHuoChatModel buildXingHuoChatModel(String apiKey) { YudaoAiProperties.XingHuo properties = new YudaoAiProperties.XingHuo() .setApiKey(apiKey).setModel(XingHuoChatModel.MODEL_DEFAULT); return new AiAutoConfiguration().buildXingHuoChatClient(properties); } /** * 可参考 {@link AiAutoConfiguration#baiChuanChatClient(YudaoAiProperties)} */ private BaiChuanChatModel buildBaiChuanChatModel(String apiKey) { YudaoAiProperties.BaiChuan properties = new YudaoAiProperties.BaiChuan() .setApiKey(apiKey); return new AiAutoConfiguration().buildBaiChuanChatClient(properties); } /** * 可参考 {@link OpenAiChatAutoConfiguration} 的 openAiChatModel 方法 */ private static OpenAiChatModel buildOpenAiChatModel(String openAiToken, String url) { return OpenAiChatModel.builder() .options(buildOpenAiChatOptions(openAiToken, url).build()) .build(); } private static OpenAiChatModel buildAzureOpenAiChatModel(String openAiToken, String url) { return OpenAiChatModel.builder() .options(buildOpenAiChatOptions(openAiToken, url) .azure(true) .build()) .build(); } private static OpenAiChatOptions.Builder buildOpenAiChatOptions(String apiKey, String url) { OpenAiChatOptions.Builder optionsBuilder = OpenAiChatOptions.builder().apiKey(apiKey); if (StrUtil.isNotEmpty(url)) { optionsBuilder.baseUrl(url); } return optionsBuilder; } /** * 可参考 {@link AnthropicChatAutoConfiguration} 的 anthropicApi 方法 */ private static AnthropicChatModel buildAnthropicChatModel(String apiKey, String url) { AnthropicChatOptions.Builder optionsBuilder = AnthropicChatOptions.builder().apiKey(apiKey); if (StrUtil.isNotEmpty(url)) { optionsBuilder.baseUrl(url); } return AnthropicChatModel.builder() .options(optionsBuilder.build()) .build(); } /** * 可参考 {@link GoogleGenAiChatAutoConfiguration} 的 googleGenAiChatModel 方法 */ private static GoogleGenAiChatModel buildGeminiChatModel(String apiKey, String url) { Client.Builder clientBuilder = Client.builder().apiKey(apiKey); if (StrUtil.isNotBlank(url)) { clientBuilder.httpOptions(HttpOptions.builder() .baseUrl(url) // TeamOrouter 的 Gemini 原生协议使用 Authorization Bearer 鉴权 .headers(Collections.singletonMap("Authorization", "Bearer " + apiKey)) .build()); } return GoogleGenAiChatModel.builder() .genAiClient(clientBuilder.build()) .options(GoogleGenAiChatOptions.builder() .model("gemini-2.5-flash") .build()) .toolCallingManager(SpringUtil.getBean(ToolCallingManager.class)) .retryTemplate(RetryUtils.DEFAULT_RETRY_TEMPLATE) .observationRegistry(SpringUtil.getBean(ObservationRegistry.class)) .build(); } /** * 可参考 {@link OpenAiImageAutoConfiguration} 的 openAiImageModel 方法 */ private OpenAiImageModel buildOpenAiImageModel(String openAiToken, String url) { OpenAiImageOptions.Builder optionsBuilder = OpenAiImageOptions.builder().apiKey(openAiToken); if (StrUtil.isNotEmpty(url)) { optionsBuilder.baseUrl(url); } return OpenAiImageModel.builder() .options(optionsBuilder.build()) .build(); } /** * 创建 SiliconFlowImageModel 对象 */ private SiliconFlowImageModel buildSiliconFlowImageModel(String apiToken, String url) { url = StrUtil.blankToDefault(url, SiliconFlowApiConstants.DEFAULT_BASE_URL); SiliconFlowImageApi openAiApi = new SiliconFlowImageApi(url, apiToken); return new SiliconFlowImageModel(openAiApi); } /** * 可参考 {@link OllamaChatAutoConfiguration} 的 ollamaChatModel 方法 */ private static OllamaChatModel buildOllamaChatModel(String url) { OllamaApi ollamaApi = OllamaApi.builder().baseUrl(url).build(); return OllamaChatModel.builder() .ollamaApi(ollamaApi) .build(); } /** * 可参考 {@link StabilityAiImageAutoConfiguration} 的 stabilityAiImageModel 方法 */ private StabilityAiImageModel buildStabilityAiImageModel(String apiKey, String url) { url = StrUtil.blankToDefault(url, StabilityAiApi.DEFAULT_BASE_URL); StabilityAiApi stabilityAiApi = new StabilityAiApi(apiKey, StabilityAiApi.DEFAULT_IMAGE_MODEL, url); return new StabilityAiImageModel(stabilityAiApi); } private ChatModel buildGrokChatModel(String apiKey,String url) { YudaoAiProperties.Grok properties = new YudaoAiProperties.Grok() .setBaseUrl(url) .setApiKey(apiKey); return new AiAutoConfiguration().buildGrokChatClient(properties); } // ========== 各种创建 EmbeddingModel 的方法 ========== /** * 可参考 {@link DashScopeEmbeddingAutoConfiguration} 的 DashScopeEmbeddingModel 方法 */ private DashScopeEmbeddingModel buildTongYiEmbeddingModel(String apiKey, String model) { return AiAutoConfiguration.buildTongYiEmbeddingModel(apiKey, model); } private OllamaEmbeddingModel buildOllamaEmbeddingModel(String url, String model) { OllamaApi ollamaApi = OllamaApi.builder().baseUrl(url).build(); OllamaEmbeddingOptions ollamaOptions = OllamaEmbeddingOptions.builder().model(model).build(); return OllamaEmbeddingModel.builder() .ollamaApi(ollamaApi) .options(ollamaOptions) .build(); } /** * 可参考 {@link OpenAiEmbeddingAutoConfiguration} 的 openAiEmbeddingModel 方法 */ private OpenAiEmbeddingModel buildOpenAiEmbeddingModel(String openAiToken, String url, String model) { OpenAiEmbeddingOptions.Builder optionsBuilder = OpenAiEmbeddingOptions.builder() .apiKey(openAiToken) .model(model); if (StrUtil.isNotEmpty(url)) { optionsBuilder.baseUrl(url); } return OpenAiEmbeddingModel.builder() .metadataMode(MetadataMode.EMBED) .options(optionsBuilder.build()) .build(); } private OpenAiEmbeddingModel buildAzureOpenAiEmbeddingModel(String openAiToken, String url, String model) { OpenAiEmbeddingOptions.Builder optionsBuilder = OpenAiEmbeddingOptions.builder() .apiKey(openAiToken) .model(model) .deploymentName(model) .azure(true); if (StrUtil.isNotEmpty(url)) { optionsBuilder.baseUrl(url); } return OpenAiEmbeddingModel.builder() .metadataMode(MetadataMode.EMBED) .options(optionsBuilder.build()) .build(); } // ========== 各种创建 VectorStore 的方法 ========== /** * 注意:仅适合本地测试使用,生产建议还是使用 Qdrant、Milvus 等 */ @SneakyThrows @SuppressWarnings("ResultOfMethodCallIgnored") private SimpleVectorStore buildSimpleVectorStore(EmbeddingModel embeddingModel) { SimpleVectorStore vectorStore = SimpleVectorStore.builder(embeddingModel).build(); // 启动加载 File file = new File(StrUtil.format("{}/vector_store/simple_{}.json", FileUtil.getUserHomePath(), embeddingModel.getClass().getSimpleName())); if (!file.exists()) { FileUtil.mkParentDirs(file); file.createNewFile(); } else if (file.length() > 0) { vectorStore.load(file); } // 定时持久化,每分钟一次 Timer timer = new Timer("SimpleVectorStoreTimer-" + file.getAbsolutePath()); timer.scheduleAtFixedRate(new TimerTask() { @Override public void run() { vectorStore.save(file); } }, Duration.ofMinutes(1).toMillis(), Duration.ofMinutes(1).toMillis()); // 关闭时,进行持久化 RuntimeUtil.addShutdownHook(() -> vectorStore.save(file)); return vectorStore; } /** * 参考 {@link QdrantVectorStoreAutoConfiguration} 的 vectorStore 方法 */ @SneakyThrows private QdrantVectorStore buildQdrantVectorStore(EmbeddingModel embeddingModel) { QdrantVectorStoreAutoConfiguration configuration = new QdrantVectorStoreAutoConfiguration(); QdrantVectorStoreProperties properties = SpringUtil.getBean(QdrantVectorStoreProperties.class); // 参考 QdrantVectorStoreAutoConfiguration 实现,创建 QdrantClient 对象 QdrantGrpcClient.Builder grpcClientBuilder = QdrantGrpcClient.newBuilder( properties.getHost(), properties.getPort(), properties.isUseTls()); if (StrUtil.isNotEmpty(properties.getApiKey())) { grpcClientBuilder.withApiKey(properties.getApiKey()); } QdrantClient qdrantClient = new QdrantClient(grpcClientBuilder.build()); // 创建 QdrantVectorStore 对象 QdrantVectorStore vectorStore = configuration.vectorStore(embeddingModel, properties, qdrantClient, getObservationRegistry(), getCustomObservationConvention(), getBatchingStrategy()); // 初始化索引 vectorStore.afterPropertiesSet(); return vectorStore; } /** * 参考 {@link RedisVectorStoreAutoConfiguration} 的 vectorStore 方法 */ private RedisVectorStore buildRedisVectorStore(EmbeddingModel embeddingModel, Map<String, Class<?>> metadataFields) { // 创建 RedisClient 对象 RedisClient redisClient = buildRedisClient(); // 创建 RedisVectorStoreProperties 对象 RedisVectorStoreProperties properties = SpringUtil.getBean(RedisVectorStoreProperties.class); RedisVectorStore redisVectorStore = RedisVectorStore.builder(redisClient, embeddingModel) .indexName(properties.getIndexName()).prefix(properties.getPrefix()) .initializeSchema(properties.isInitializeSchema()) .metadataFields(convertList(metadataFields.entrySet(), entry -> { String fieldName = entry.getKey(); Class<?> fieldType = entry.getValue(); if (Number.class.isAssignableFrom(fieldType)) { return RedisVectorStore.MetadataField.numeric(fieldName); } if (Boolean.class.isAssignableFrom(fieldType)) { return RedisVectorStore.MetadataField.tag(fieldName); } return RedisVectorStore.MetadataField.text(fieldName); })) .observationRegistry(getObservationRegistry().getObject()) .customObservationConvention(getCustomObservationConvention().getObject()) .batchingStrategy(getBatchingStrategy()) .build(); // 初始化索引 redisVectorStore.afterPropertiesSet(); return redisVectorStore; } private RedisClient buildRedisClient() { DataRedisProperties redisProperties = SpringUtil.getBean(DataRedisProperties.class); Assert.isNull(redisProperties.getCluster(), "RedisVectorStore 暂不支持 Redis Cluster 模式"); Assert.isNull(redisProperties.getSentinel(), "RedisVectorStore 暂不支持 Redis Sentinel 模式"); Assert.isNull(redisProperties.getMasterreplica(), "RedisVectorStore 暂不支持 Redis Master-Replica 模式"); if (StrUtil.isNotEmpty(redisProperties.getUrl())) { return RedisClient.create(redisProperties.getUrl()); } DefaultJedisClientConfig.Builder clientConfigBuilder = DefaultJedisClientConfig.builder() .ssl(redisProperties.getSsl().isEnabled()) .database(redisProperties.getDatabase()); if (StrUtil.isNotEmpty(redisProperties.getUsername())) { clientConfigBuilder.user(redisProperties.getUsername()); } if (StrUtil.isNotEmpty(redisProperties.getPassword())) { clientConfigBuilder.password(redisProperties.getPassword()); } if (StrUtil.isNotEmpty(redisProperties.getClientName())) { clientConfigBuilder.clientName(redisProperties.getClientName()); } if (redisProperties.getTimeout() != null) { clientConfigBuilder.socketTimeoutMillis(toMillis(redisProperties.getTimeout())); } if (redisProperties.getConnectTimeout() != null) { clientConfigBuilder.connectionTimeoutMillis(toMillis(redisProperties.getConnectTimeout())); } JedisClientConfig clientConfig = clientConfigBuilder.build(); return RedisClient.builder() .hostAndPort(new HostAndPort(redisProperties.getHost(), redisProperties.getPort())) .clientConfig(clientConfig) .build(); } private static int toMillis(Duration duration) { return Math.toIntExact(duration.toMillis()); } /** * 参考 {@link MilvusVectorStoreAutoConfiguration} 的 vectorStore 方法 */ @SneakyThrows private MilvusVectorStore buildMilvusVectorStore(EmbeddingModel embeddingModel) { MilvusVectorStoreAutoConfiguration configuration = new MilvusVectorStoreAutoConfiguration(); // 获取配置属性 MilvusVectorStoreProperties serverProperties = SpringUtil.getBean(MilvusVectorStoreProperties.class); MilvusServiceClientProperties clientProperties = SpringUtil.getBean(MilvusServiceClientProperties.class); // 创建 MilvusServiceClient 对象 MilvusServiceClient milvusClient = configuration.milvusClient(serverProperties, clientProperties, new MilvusServiceClientConnectionDetails() { var connectParam = io.milvus.param.ConnectParam.newBuilder() .withHost(clientProperties.getHost()) .withPort(clientProperties.getPort()) .withDatabaseName(serverProperties.getDatabaseName()) .build(); var milvusClient = new io.milvus.client.MilvusServiceClient(connectParam); @Override public String getHost() { return clientProperties.getHost(); } @Override public int getPort() { return clientProperties.getPort(); } } ); // 创建 MilvusVectorStore 对象 MilvusVectorStore vectorStore = configuration.vectorStore(milvusClient, embeddingModel, serverProperties, getBatchingStrategy(), getObservationRegistry(), getCustomObservationConvention()); // 初始化索引 vectorStore.afterPropertiesSet(); MilvusVectorStore vectorStore = MilvusVectorStore.builder(milvusClient, embeddingModel) .databaseName(serverProperties.getDatabaseName()) .collectionName(serverProperties.getCollectionName()) .initializeSchema(serverProperties.isInitializeSchema()) .batchingStrategy(new TokenCountBatchingStrategy()) .build(); try { vectorStore.afterPropertiesSet(); } catch (Exception e) { throw new RuntimeException("Milvus 向量存储初始化失败: " + e.getMessage(), e); } return vectorStore; } private static ObjectProvider<ObservationRegistry> getObservationRegistry() { return new ObjectProvider<>() { @Override public ObservationRegistry getObject() throws BeansException { return SpringUtil.getBean(ObservationRegistry.class); } }; } private static ObjectProvider<VectorStoreObservationConvention> getCustomObservationConvention() { return new ObjectProvider<>() { @Override public VectorStoreObservationConvention getObject() throws BeansException { return new DefaultVectorStoreObservationConvention(); } }; } private static BatchingStrategy getBatchingStrategy() { return SpringUtil.getBean(BatchingStrategy.class); } private static ObjectProvider<EmbeddingModelObservationConvention> getEmbeddingModelObservationConvention() { return new ObjectProvider<>() { @Override public EmbeddingModelObservationConvention getObject() throws BeansException { return SpringUtil.getBean(EmbeddingModelObservationConvention.class); } }; private static String buildCacheKey(Class<?> clazz, Object... params) { if (ArrayUtil.isEmpty(params)) return clazz.getName(); return StrUtil.format("{}#{}", clazz.getName(), ArrayUtil.join(params, "_")); } } yudao-module-ai/src/main/java/cn/iocoder/yudao/module/ai/service/knowledge/AiKnowledgeDocumentServiceImpl.java
@@ -1,226 +1,363 @@ package cn.iocoder.yudao.module.ai.service.knowledge; import cn.hutool.core.collection.CollUtil; import cn.hutool.core.util.ObjUtil; import cn.hutool.core.io.FileUtil; import cn.hutool.core.io.IoUtil; import cn.hutool.core.util.ArrayUtil; import cn.hutool.core.util.StrUtil; import cn.hutool.http.HttpUtil; import cn.iocoder.yudao.framework.common.enums.CommonStatusEnum; import cn.iocoder.yudao.framework.common.pojo.PageResult; import cn.iocoder.yudao.framework.common.util.object.BeanUtils; import cn.iocoder.yudao.module.ai.controller.admin.knowledge.vo.document.AiKnowledgeDocumentCreateListReqVO; import cn.iocoder.yudao.module.ai.controller.admin.knowledge.vo.document.AiKnowledgeDocumentPageReqVO; import cn.iocoder.yudao.module.ai.controller.admin.knowledge.vo.document.AiKnowledgeDocumentUpdateReqVO; import cn.iocoder.yudao.module.ai.controller.admin.knowledge.vo.document.AiKnowledgeDocumentUpdateStatusReqVO; import cn.iocoder.yudao.module.ai.controller.admin.knowledge.vo.knowledge.AiKnowledgeDocumentCreateReqVO; import cn.iocoder.yudao.module.ai.controller.admin.knowledge.vo.document.*; import cn.iocoder.yudao.module.ai.controller.admin.knowledge.vo.segment.AiKnowledgeSegmentProcessRespVO; import cn.iocoder.yudao.module.ai.dal.dataobject.knowledge.AiKnowledgeDO; import cn.iocoder.yudao.module.ai.dal.dataobject.knowledge.AiKnowledgeDocumentDO; import cn.iocoder.yudao.module.ai.dal.dataobject.knowledge.AiKnowledgeSegmentDO; import cn.iocoder.yudao.module.ai.dal.dataobject.model.AiModelDO; import cn.iocoder.yudao.module.ai.dal.mysql.knowledge.AiKnowledgeDocumentMapper; import cn.iocoder.yudao.module.ai.dal.mysql.knowledge.AiKnowledgeSegmentMapper; import cn.iocoder.yudao.module.ai.service.model.AiModelService; import jakarta.annotation.Resource; import lombok.extern.slf4j.Slf4j; import org.springframework.ai.document.Document; import org.springframework.ai.reader.tika.TikaDocumentReader; import org.springframework.ai.tokenizer.TokenCountEstimator; import org.springframework.context.annotation.Lazy; import org.springframework.core.io.ByteArrayResource; import org.springframework.ai.transformer.splitter.TokenTextSplitter; import org.springframework.stereotype.Service; import org.springframework.transaction.annotation.Transactional; import org.apache.tika.parser.AutoDetectParser; import org.apache.tika.parser.ParseContext; import org.apache.tika.parser.Parser; import org.apache.tika.extractor.EmbeddedDocumentExtractor; import org.apache.tika.sax.BodyContentHandler; import org.xml.sax.ContentHandler; import java.io.ByteArrayOutputStream; import java.io.InputStream; import java.net.HttpURLConnection; import java.net.URI; import java.net.URL; import java.net.URLEncoder; import java.nio.charset.StandardCharsets; import java.util.ArrayList; import java.util.Collection; import java.util.Collections; import java.util.List; import static cn.iocoder.yudao.framework.common.exception.util.ServiceExceptionUtil.exception; import static cn.iocoder.yudao.framework.common.util.collection.CollectionUtils.convertList; import static cn.iocoder.yudao.module.ai.enums.ErrorCodeConstants.*; /** * AI 知识库文档 Service 实现类 * * @author xiaoxin */ @Service @Slf4j @Service public class AiKnowledgeDocumentServiceImpl implements AiKnowledgeDocumentService { @Resource private AiKnowledgeDocumentMapper knowledgeDocumentMapper; private AiKnowledgeDocumentMapper documentMapper; @Resource private TokenCountEstimator tokenCountEstimator; private AiKnowledgeSegmentMapper segmentMapper; @Resource private AiKnowledgeSegmentService knowledgeSegmentService; @Resource @Lazy // 延迟加载,避免循环依赖 private AiKnowledgeService knowledgeService; @Override public Long createKnowledgeDocument(AiKnowledgeDocumentCreateReqVO createReqVO) { // 1. 校验参数 knowledgeService.validateKnowledgeExists(createReqVO.getKnowledgeId()); // 2. 下载文档 String content = readUrl(createReqVO.getUrl()); // 3. 文档记录入库 AiKnowledgeDocumentDO documentDO = BeanUtils.toBean(createReqVO, AiKnowledgeDocumentDO.class) .setContent(content).setContentLength(content.length()).setTokens(tokenCountEstimator.estimate(content)) .setStatus(CommonStatusEnum.ENABLE.getStatus()); knowledgeDocumentMapper.insert(documentDO); // 4. 文档切片入库(异步) knowledgeSegmentService.createKnowledgeSegmentBySplitContentAsync(documentDO.getId(), content); return documentDO.getId(); } @Override public List<Long> createKnowledgeDocumentList(AiKnowledgeDocumentCreateListReqVO createListReqVO) { // 1. 校验参数 knowledgeService.validateKnowledgeExists(createListReqVO.getKnowledgeId()); // 2. 下载文档 List<String> contents = convertList(createListReqVO.getList(), document -> readUrl(document.getUrl())); // 3. 文档记录入库 List<AiKnowledgeDocumentDO> documentDOs = new ArrayList<>(createListReqVO.getList().size()); for (int i = 0; i < createListReqVO.getList().size(); i++) { AiKnowledgeDocumentCreateListReqVO.Document documentVO = createListReqVO.getList().get(i); String content = contents.get(i); documentDOs.add(BeanUtils.toBean(documentVO, AiKnowledgeDocumentDO.class) .setKnowledgeId(createListReqVO.getKnowledgeId()) .setContent(content).setContentLength(content.length()) .setTokens(tokenCountEstimator.estimate(content)) .setSegmentMaxTokens(createListReqVO.getSegmentMaxTokens()) .setStatus(CommonStatusEnum.ENABLE.getStatus())); } knowledgeDocumentMapper.insertBatch(documentDOs); // 4. 批量创建文档切片(异步) documentDOs.forEach(documentDO -> knowledgeSegmentService .createKnowledgeSegmentBySplitContentAsync(documentDO.getId(), documentDO.getContent())); return convertList(documentDOs, AiKnowledgeDocumentDO::getId); } @Override public PageResult<AiKnowledgeDocumentDO> getKnowledgeDocumentPage(AiKnowledgeDocumentPageReqVO pageReqVO) { return knowledgeDocumentMapper.selectPage(pageReqVO); } @Override public AiKnowledgeDocumentDO getKnowledgeDocument(Long id) { return knowledgeDocumentMapper.selectById(id); } @Override public void updateKnowledgeDocument(AiKnowledgeDocumentUpdateReqVO reqVO) { // 1. 校验文档是否存在 AiKnowledgeDocumentDO oldDocument = validateKnowledgeDocumentExists(reqVO.getId()); // 2. 更新文档 AiKnowledgeDocumentDO document = BeanUtils.toBean(reqVO, AiKnowledgeDocumentDO.class); knowledgeDocumentMapper.updateById(document); // 3. 如果处于开启状态,并且最大 tokens 发生变化,则 segment 需要重新索引 if (CommonStatusEnum.isEnable(oldDocument.getStatus()) && reqVO.getSegmentMaxTokens() != null && ObjUtil.notEqual(reqVO.getSegmentMaxTokens(), oldDocument.getSegmentMaxTokens())) { // 删除旧的文档切片 knowledgeSegmentService.deleteKnowledgeSegmentByDocumentId(reqVO.getId()); // 重新创建文档切片 knowledgeSegmentService.createKnowledgeSegmentBySplitContentAsync(reqVO.getId(), oldDocument.getContent()); } } @Override public void updateKnowledgeDocumentStatus(AiKnowledgeDocumentUpdateStatusReqVO reqVO) { // 1. 校验存在 AiKnowledgeDocumentDO document = validateKnowledgeDocumentExists(reqVO.getId()); // 2. 更新状态 knowledgeDocumentMapper.updateById(new AiKnowledgeDocumentDO() .setId(reqVO.getId()).setStatus(reqVO.getStatus())); // 3. 处理文档切片 if (CommonStatusEnum.isEnable(reqVO.getStatus())) { knowledgeSegmentService.createKnowledgeSegmentBySplitContentAsync(reqVO.getId(), document.getContent()); } else { knowledgeSegmentService.deleteKnowledgeSegmentByDocumentId(reqVO.getId()); } } @Resource private AiKnowledgeSegmentService segmentService; @Resource private AiModelService modelService; @Override @Transactional(rollbackFor = Exception.class) public void deleteKnowledgeDocument(Long id) { // 1. 校验存在 validateKnowledgeDocumentExists(id); // 2. 删除 knowledgeDocumentMapper.deleteById(id); // 3. 删除对应的段落 knowledgeSegmentService.deleteKnowledgeSegmentByDocumentId(id); } @Override public AiKnowledgeDocumentDO validateKnowledgeDocumentExists(Long id) { AiKnowledgeDocumentDO knowledgeDocument = knowledgeDocumentMapper.selectById(id); if (knowledgeDocument == null) { throw exception(KNOWLEDGE_DOCUMENT_NOT_EXISTS); public List<Long> createDocuments(AiKnowledgeDocumentCreateReqVO createReqVO) { AiKnowledgeDO knowledge = knowledgeService.validateKnowledge(createReqVO.getKnowledgeId()); List<AiKnowledgeDocumentCreateReqVO.DocumentItem> list = createReqVO.getList(); if (CollUtil.isEmpty(list)) { throw new IllegalArgumentException("文档列表不能为空"); } return knowledgeDocument; } int defaultSegmentMaxTokens = createReqVO.getSegmentMaxTokens() != null ? createReqVO.getSegmentMaxTokens() : 800; @Override public String readUrl(String url) { // 下载文件 ByteArrayResource resource; try { byte[] bytes = HttpUtil.downloadBytes(url); if (bytes.length == 0) { throw exception(KNOWLEDGE_DOCUMENT_FILE_EMPTY); List<Long> ids = new ArrayList<>(); for (AiKnowledgeDocumentCreateReqVO.DocumentItem item : list) { AiKnowledgeDocumentCreateReqVO.UrlInfo urlInfo = item.getUrl(); String fileUrl = urlInfo != null ? urlInfo.getUrl() : null; AiKnowledgeDocumentDO document = new AiKnowledgeDocumentDO(); document.setKnowledgeId(knowledge.getId()); document.setName(item.getName()); document.setUrl(fileUrl); document.setStatus(CommonStatusEnum.ENABLE.getStatus()); document.setSegmentMaxTokens(defaultSegmentMaxTokens); documentMapper.insert(document); ids.add(document.getId()); if (fileUrl != null) { try { loadDocumentContent(document); documentMapper.updateById(document); // 自动执行分段和向量化 processDocumentSegmentsInternal(document, knowledge); } catch (Exception e) { log.error("文档[{}]处理失败: {}", document.getId(), e.getMessage()); document.setStatus(CommonStatusEnum.DISABLE.getStatus()); documentMapper.updateById(document); } } resource = new ByteArrayResource(bytes); } return ids; } private void loadDocumentContent(AiKnowledgeDocumentDO document) { try { byte[] fileBytes = downloadFile(document.getUrl()); String content = extractText(fileBytes, document.getName()); if (StrUtil.isEmpty(content)) { content = new String(fileBytes, java.nio.charset.Charset.forName("UTF-8")); } document.setContent(content); document.setContentLength(content.length()); document.setTokens(estimateTokens(content)); } catch (Exception e) { log.error("[readUrl][url({}) 读取失败]", url, e); throw exception(KNOWLEDGE_DOCUMENT_FILE_DOWNLOAD_FAIL); throw new RuntimeException("文档内容加载失败: " + e.getMessage(), e); } // 读取文件 TikaDocumentReader loader = new TikaDocumentReader(resource); List<Document> documents = loader.get(); Document document = CollUtil.getFirst(documents); if (document == null || StrUtil.isEmpty(document.getText())) { throw exception(KNOWLEDGE_DOCUMENT_FILE_READ_FAIL); } return document.getText(); } @Override public List<AiKnowledgeDocumentDO> getKnowledgeDocumentList(Collection<Long> ids) { if (CollUtil.isEmpty(ids)) { return Collections.emptyList(); private byte[] downloadFile(String fileUrl) { try { // 手动编码 URL 路径中的非 ASCII 字符,避免 URI 构造函数报错 String encodedUrl = encodeUrlPath(fileUrl); URL url = URL.of(new URI(encodedUrl), null); HttpURLConnection conn = (HttpURLConnection) url.openConnection(); conn.setConnectTimeout(30000); conn.setReadTimeout(60000); conn.setRequestMethod("GET"); try (InputStream is = conn.getInputStream(); ByteArrayOutputStream bos = new ByteArrayOutputStream()) { IoUtil.copy(is, bos); return bos.toByteArray(); } } catch (Exception e) { throw new RuntimeException("文件下载失败: " + e.getMessage(), e); } return knowledgeDocumentMapper.selectByIds(ids); } @Override public List<AiKnowledgeDocumentDO> getKnowledgeDocumentListByKnowledgeId(Long knowledgeId) { return knowledgeDocumentMapper.selectListByKnowledgeId(knowledgeId); private static String encodeUrlPath(String urlString) { // 分离 scheme://authority 和 path?query#fragment int schemeEnd = urlString.indexOf("://"); if (schemeEnd < 0) return urlString; int pathStart = urlString.indexOf('/', schemeEnd + 3); if (pathStart < 0) return urlString; // 无 path String base = urlString.substring(0, pathStart); String pathAndRest = urlString.substring(pathStart); // 分离 path 和 query + fragment int queryStart = pathAndRest.indexOf('?'); int fragStart = pathAndRest.indexOf('#'); String rawPath, query, fragment; if (queryStart >= 0) { rawPath = pathAndRest.substring(0, queryStart); if (fragStart >= 0 && fragStart > queryStart) { query = pathAndRest.substring(queryStart, fragStart); fragment = pathAndRest.substring(fragStart); } else { query = pathAndRest.substring(queryStart); fragment = ""; } } else if (fragStart >= 0) { rawPath = pathAndRest.substring(0, fragStart); query = ""; fragment = pathAndRest.substring(fragStart); } else { rawPath = pathAndRest; query = ""; fragment = ""; } // 对路径每个段做 URL 编码 StringBuilder encodedPath = new StringBuilder(); for (String segment : rawPath.split("/")) { if (!segment.isEmpty()) { encodedPath.append("/").append(URLEncoder.encode(segment, StandardCharsets.UTF_8) .replace("+", "%20")); } } if (rawPath.endsWith("/")) encodedPath.append("/"); return base + encodedPath + query + fragment; } private String extractText(byte[] fileBytes, String fileName) { try { String ext = FileUtil.extName(fileName).toLowerCase(); if (ArrayUtil.contains(new String[]{"txt", "md", "json", "xml", "csv", "yaml", "yml"}, ext)) { return new String(fileBytes, StandardCharsets.UTF_8); } // 使用 AutoDetectParser + EmbeddedDocumentExtractor 跳过嵌入图片,避免提取到二进制乱码 Parser parser = new AutoDetectParser(); ParseContext context = new ParseContext(); context.set(EmbeddedDocumentExtractor.class, new EmbeddedDocumentExtractor() { @Override public boolean shouldParseEmbedded(org.apache.tika.metadata.Metadata metadata) { return false; } @Override public void parseEmbedded(InputStream inputStream, ContentHandler contentHandler, org.apache.tika.metadata.Metadata metadata, boolean outputHtml) { } }); BodyContentHandler handler = new BodyContentHandler(-1); try (InputStream is = new java.io.ByteArrayInputStream(fileBytes)) { parser.parse(is, handler, new org.apache.tika.metadata.Metadata(), context); } return handler.toString().trim(); } catch (Exception e) { log.warn("Tika 解析文档失败,尝试 UTF-8 文本读取: {}", e.getMessage()); return ""; } } @Override @Transactional(rollbackFor = Exception.class) public void deleteKnowledgeDocumentByKnowledgeId(Long knowledgeId) { // 1. 获取该知识库下的所有文档 List<AiKnowledgeDocumentDO> documents = knowledgeDocumentMapper.selectListByKnowledgeId(knowledgeId); if (CollUtil.isEmpty(documents)) { return; public void updateDocument(AiKnowledgeDocumentUpdateReqVO updateReqVO) { AiKnowledgeDocumentDO document = validateDocumentExists(updateReqVO.getId()); boolean urlChanged = StrUtil.isNotEmpty(updateReqVO.getUrl()) && !updateReqVO.getUrl().equals(document.getUrl()); if (StrUtil.isNotEmpty(updateReqVO.getName())) document.setName(updateReqVO.getName()); if (urlChanged) document.setUrl(updateReqVO.getUrl()); documentMapper.updateById(document); // URL 变更时重新提取内容、重新分段和向量化 if (urlChanged) { try { loadDocumentContent(document); documentMapper.updateById(document); AiKnowledgeDO knowledge = knowledgeService.validateKnowledge(document.getKnowledgeId()); // 删除旧分段和向量 segmentService.deleteSegmentsByDocumentId(document.getId()); // 重新分段 + 向量化 processDocumentSegmentsInternal(document, knowledge); } catch (Exception e) { log.error("文档[{}] URL 更新后处理失败: {}", document.getId(), e.getMessage()); document.setStatus(CommonStatusEnum.DISABLE.getStatus()); documentMapper.updateById(document); } } } // 2. 逐个删除文档及其对应的段落 for (AiKnowledgeDocumentDO document : documents) { deleteKnowledgeDocument(document.getId()); @Override @Transactional(rollbackFor = Exception.class) public void deleteDocument(Long id) { AiKnowledgeDocumentDO document = validateDocumentExists(id); segmentService.deleteSegmentsByDocumentId(id); documentMapper.deleteById(id); } @Override public AiKnowledgeDocumentDO getDocument(Long id) { return documentMapper.selectById(id); } @Override public AiKnowledgeDocumentDO validateDocumentExists(Long id) { AiKnowledgeDocumentDO document = documentMapper.selectById(id); if (document == null) throw exception(KNOWLEDGE_DOCUMENT_NOT_EXISTS); return document; } @Override public PageResult<AiKnowledgeDocumentDO> getDocumentPage(AiKnowledgeDocumentPageReqVO pageReqVO) { return documentMapper.selectPage(pageReqVO); } @Override @Transactional(rollbackFor = Exception.class) public void updateDocumentStatus(AiKnowledgeDocumentUpdateStatusReqVO updateStatusReqVO) { AiKnowledgeDocumentDO document = validateDocumentExists(updateStatusReqVO.getId()); document.setStatus(updateStatusReqVO.getStatus()); documentMapper.updateById(document); } @Override public List<AiKnowledgeSegmentProcessRespVO> getDocumentProcessingProgress(List<Long> documentIds) { if (CollUtil.isEmpty(documentIds)) return Collections.emptyList(); return segmentMapper.selectProcessList(documentIds); } @Override public void processDocumentSegments(Long documentId) { AiKnowledgeDocumentDO document = validateDocumentExists(documentId); if (StrUtil.isEmpty(document.getContent())) { throw exception(KNOWLEDGE_DOCUMENT_FILE_EMPTY); } AiKnowledgeDO knowledge = knowledgeService.validateKnowledge(document.getKnowledgeId()); processDocumentSegmentsInternal(document, knowledge); } private void processDocumentSegmentsInternal(AiKnowledgeDocumentDO document, AiKnowledgeDO knowledge) { if (StrUtil.isEmpty(document.getContent())) return; AiModelDO embeddingModel = modelService.validateModel(knowledge.getEmbeddingModelId()); // 删除旧分段 segmentService.deleteSegmentsByDocumentId(document.getId()); // 文本切片 int segmentMaxTokens = document.getSegmentMaxTokens() != null ? document.getSegmentMaxTokens() : 800; List<String> segmentTexts = splitContent(document.getContent(), segmentMaxTokens); if (CollUtil.isEmpty(segmentTexts)) return; // 向量化并存储 List<AiKnowledgeSegmentDO> segments = new ArrayList<>(); for (String content : segmentTexts) { if (StrUtil.isEmpty(content.trim())) continue; segments.add(buildSegmentDO(knowledge.getId(), document.getId(), content)); } segmentService.saveSegments(segments, knowledge.getId(), embeddingModel.getId()); } private List<String> splitContent(String content, int maxTokens) { TokenTextSplitter splitter = TokenTextSplitter.builder() .withChunkSize(maxTokens) .withMinChunkSizeChars(50) .withMinChunkLengthToEmbed(10) .withMaxNumChunks(1000) .withKeepSeparator(true) .build(); List<Document> docs = splitter.apply(Collections.singletonList(new Document(content))); List<String> result = new ArrayList<>(); for (Document doc : docs) { if (StrUtil.isNotEmpty(doc.getText())) { result.add(doc.getText().trim()); } } return result; } @Override public List<String> previewSplit(String url, String name, Integer segmentMaxTokens) { byte[] fileBytes = downloadFile(url); String text = extractText(fileBytes, name); if (StrUtil.isEmpty(text)) { text = new String(fileBytes, java.nio.charset.Charset.forName("UTF-8")); } int maxTokens = segmentMaxTokens != null && segmentMaxTokens > 0 ? segmentMaxTokens : 800; return splitContent(text, maxTokens); } private AiKnowledgeSegmentDO buildSegmentDO(Long knowledgeId, Long documentId, String content) { AiKnowledgeSegmentDO segment = new AiKnowledgeSegmentDO(); segment.setKnowledgeId(knowledgeId); segment.setDocumentId(documentId); segment.setContent(content); segment.setContentLength(content.length()); segment.setTokens(estimateTokens(content)); segment.setStatus(CommonStatusEnum.ENABLE.getStatus()); segment.setVectorId(AiKnowledgeSegmentDO.VECTOR_ID_EMPTY); segment.setRetrievalCount(0); return segment; } private Integer estimateTokens(String text) { if (StrUtil.isEmpty(text)) return 0; int chineseChars = 0, englishWords = 0; for (char c : text.toCharArray()) { if (c >= 0x4E00 && c <= 0x9FA5) chineseChars++; } for (String word : text.split("\\s+")) { if (word.matches(".*[a-zA-Z].*")) englishWords++; } return chineseChars + (int) (englishWords * 1.3); } } yudao-module-ai/src/main/java/cn/iocoder/yudao/module/ai/service/knowledge/AiKnowledgeSegmentServiceImpl.java
@@ -1,509 +1,353 @@ package cn.iocoder.yudao.module.ai.service.knowledge; import cn.hutool.core.collection.CollUtil; import cn.hutool.core.collection.ListUtil; import cn.hutool.core.util.ObjUtil; import cn.hutool.core.map.MapUtil; import cn.hutool.core.util.StrUtil; import cn.iocoder.yudao.framework.common.enums.CommonStatusEnum; import cn.iocoder.yudao.framework.common.pojo.PageResult; import cn.iocoder.yudao.framework.common.util.object.BeanUtils; import cn.iocoder.yudao.module.ai.controller.admin.knowledge.vo.segment.AiKnowledgeSegmentPageReqVO; import cn.iocoder.yudao.module.ai.controller.admin.knowledge.vo.segment.AiKnowledgeSegmentProcessRespVO; import cn.iocoder.yudao.module.ai.controller.admin.knowledge.vo.segment.AiKnowledgeSegmentSaveReqVO; import cn.iocoder.yudao.module.ai.controller.admin.knowledge.vo.segment.AiKnowledgeSegmentUpdateStatusReqVO; import cn.iocoder.yudao.module.ai.controller.admin.knowledge.vo.segment.*; import cn.iocoder.yudao.module.ai.dal.dataobject.knowledge.AiKnowledgeDO; import cn.iocoder.yudao.module.ai.dal.dataobject.knowledge.AiKnowledgeDocumentDO; import cn.iocoder.yudao.module.ai.dal.dataobject.knowledge.AiKnowledgeSegmentDO; import cn.iocoder.yudao.module.ai.dal.mysql.knowledge.AiKnowledgeDocumentMapper; import cn.iocoder.yudao.module.ai.dal.mysql.knowledge.AiKnowledgeSegmentMapper; import cn.iocoder.yudao.module.ai.enums.AiDocumentSplitStrategyEnum; import cn.iocoder.yudao.module.ai.service.knowledge.bo.AiKnowledgeSegmentSearchReqBO; import cn.iocoder.yudao.module.ai.service.knowledge.bo.AiKnowledgeSegmentSearchRespBO; import cn.iocoder.yudao.module.ai.service.knowledge.splitter.MarkdownQaSplitter; import cn.iocoder.yudao.module.ai.service.knowledge.splitter.SemanticTextSplitter; import cn.iocoder.yudao.module.ai.service.model.AiModelService; import com.alibaba.cloud.ai.dashscope.rerank.DashScopeRerankOptions; import com.alibaba.cloud.ai.model.RerankModel; import com.alibaba.cloud.ai.model.RerankRequest; import com.alibaba.cloud.ai.model.RerankResponse; import jakarta.annotation.Resource; import lombok.extern.slf4j.Slf4j; import org.springframework.ai.document.Document; import org.springframework.ai.tokenizer.TokenCountEstimator; import org.springframework.ai.transformer.splitter.TextSplitter; import org.springframework.ai.transformer.splitter.TokenTextSplitter; import org.springframework.ai.vectorstore.SearchRequest; import org.springframework.ai.vectorstore.VectorStore; import org.springframework.ai.vectorstore.filter.Filter; import org.springframework.ai.vectorstore.filter.FilterExpressionBuilder; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.context.annotation.Lazy; import org.springframework.stereotype.Service; import org.springframework.transaction.annotation.Transactional; import java.util.*; import java.util.stream.Collectors; import static cn.iocoder.yudao.framework.common.exception.util.ServiceExceptionUtil.exception; import static cn.iocoder.yudao.framework.common.util.collection.CollectionUtils.convertList; import static cn.iocoder.yudao.module.ai.enums.ErrorCodeConstants.*; import static org.springframework.ai.vectorstore.SearchRequest.SIMILARITY_THRESHOLD_ACCEPT_ALL; /** * AI 知识库分片 Service 实现类 * * @author xiaoxin */ @Service @Slf4j @Service public class AiKnowledgeSegmentServiceImpl implements AiKnowledgeSegmentService { private static final String VECTOR_STORE_METADATA_KNOWLEDGE_ID = "knowledgeId"; private static final String VECTOR_STORE_METADATA_DOCUMENT_ID = "documentId"; private static final String VECTOR_STORE_METADATA_SEGMENT_ID = "segmentId"; private static final Map<String, Class<?>> VECTOR_STORE_METADATA_TYPES = Map.of( VECTOR_STORE_METADATA_KNOWLEDGE_ID, String.class, VECTOR_STORE_METADATA_DOCUMENT_ID, String.class, VECTOR_STORE_METADATA_SEGMENT_ID, String.class); /** * Rerank 在向量检索时,检索数量 * 该系数,目的是为了提升 Rerank 的效果 */ private static final Integer RERANK_RETRIEVAL_FACTOR = 4; private static final String METADATA_KNOWLEDGE_ID = "knowledgeId"; private static final String METADATA_DOCUMENT_ID = "documentId"; private static final String METADATA_SEGMENT_ID = "segmentId"; @Resource private AiKnowledgeSegmentMapper segmentMapper; @Resource private AiKnowledgeDocumentMapper documentMapper; @Resource private AiKnowledgeService knowledgeService; @Resource @Lazy // 延迟加载,避免循环依赖 private AiKnowledgeDocumentService knowledgeDocumentService; @Resource private AiModelService modelService; @Resource private TokenCountEstimator tokenCountEstimator; @Autowired(required = false) // 由于 spring.ai.model.rerank 配置项,可以关闭 RerankModel 的功能,所以这里只能不强制注入 private RerankModel rerankModel; @Override public PageResult<AiKnowledgeSegmentDO> getKnowledgeSegmentPage(AiKnowledgeSegmentPageReqVO pageReqVO) { return segmentMapper.selectPage(pageReqVO); } @Override public void createKnowledgeSegmentBySplitContent(Long documentId, String content) { // 1. 校验 AiKnowledgeDocumentDO documentDO = knowledgeDocumentService.validateKnowledgeDocumentExists(documentId); AiKnowledgeDO knowledgeDO = knowledgeService.validateKnowledgeExists(documentDO.getKnowledgeId()); VectorStore vectorStore = getVectorStoreById(knowledgeDO); // 2. 文档切片(使用自动检测策略) List<Document> documentSegments = splitContentByStrategy(content, documentDO.getSegmentMaxTokens(), AiDocumentSplitStrategyEnum.AUTO, documentDO.getUrl()); // 3.1 存储切片 List<AiKnowledgeSegmentDO> segmentDOs = convertList(documentSegments, segment -> { if (StrUtil.isEmpty(segment.getText())) { return null; } return new AiKnowledgeSegmentDO().setKnowledgeId(documentDO.getKnowledgeId()).setDocumentId(documentId) .setContent(segment.getText()).setContentLength(segment.getText().length()) .setVectorId(AiKnowledgeSegmentDO.VECTOR_ID_EMPTY) .setTokens(tokenCountEstimator.estimate(segment.getText())) .setStatus(CommonStatusEnum.ENABLE.getStatus()); }); segmentMapper.insertBatch(segmentDOs); // 3.2 切片向量化 for (int i = 0; i < documentSegments.size(); i++) { Document segment = documentSegments.get(i); AiKnowledgeSegmentDO segmentDO = segmentDOs.get(i); writeVectorStore(vectorStore, segmentDO, segment); } } @Override public void updateKnowledgeSegment(AiKnowledgeSegmentSaveReqVO reqVO) { // 1. 校验 AiKnowledgeSegmentDO oldSegment = validateKnowledgeSegmentExists(reqVO.getId()); // 2. 删除向量 VectorStore vectorStore = getVectorStoreById(oldSegment.getKnowledgeId()); deleteVectorStore(vectorStore, oldSegment); // 3.1 更新切片 AiKnowledgeSegmentDO newSegment = BeanUtils.toBean(reqVO, AiKnowledgeSegmentDO.class); segmentMapper.updateById(newSegment); // 3.2 重新向量化,必须开启状态 if (CommonStatusEnum.isEnable(oldSegment.getStatus())) { newSegment.setKnowledgeId(oldSegment.getKnowledgeId()).setDocumentId(oldSegment.getDocumentId()); writeVectorStore(vectorStore, newSegment, new Document(newSegment.getContent())); } } @Override public void deleteKnowledgeSegment(Long id) { // 1. 校验段落存在 AiKnowledgeSegmentDO segment = validateKnowledgeSegmentExists(id); // 2. 删除向量 VectorStore vectorStore = getVectorStoreById(segment.getKnowledgeId()); deleteVectorStore(vectorStore, segment); // 3. 删除段落记录 segmentMapper.deleteById(id); } @Override public void deleteKnowledgeSegmentByDocumentId(Long documentId) { // 1. 查询需要删除的段落 List<AiKnowledgeSegmentDO> segments = segmentMapper.selectListByDocumentId(documentId); if (CollUtil.isEmpty(segments)) { return; } // 2. 批量删除段落记录 segmentMapper.deleteByIds(convertList(segments, AiKnowledgeSegmentDO::getId)); // 3. 删除向量存储中的段落 VectorStore vectorStore = getVectorStoreById(segments.getFirst().getKnowledgeId()); vectorStore.delete(convertList(segments, AiKnowledgeSegmentDO::getVectorId)); } @Override public void updateKnowledgeSegmentStatus(AiKnowledgeSegmentUpdateStatusReqVO reqVO) { // 1. 校验 AiKnowledgeSegmentDO segment = validateKnowledgeSegmentExists(reqVO.getId()); // 2. 获取知识库向量实例 VectorStore vectorStore = getVectorStoreById(segment.getKnowledgeId()); // 3. 更新状态 segmentMapper.updateById(new AiKnowledgeSegmentDO().setId(reqVO.getId()).setStatus(reqVO.getStatus())); // 4. 更新向量 if (CommonStatusEnum.isEnable(reqVO.getStatus())) { writeVectorStore(vectorStore, segment, new Document(segment.getContent())); } else { deleteVectorStore(vectorStore, segment); } } @Override public void reindexKnowledgeSegmentByKnowledgeId(Long knowledgeId) { // 1.1 校验知识库存在 AiKnowledgeDO knowledge = knowledgeService.validateKnowledgeExists(knowledgeId); // 1.2 获取知识库向量实例 VectorStore vectorStore = getVectorStoreById(knowledge); // 2.1 查询知识库下的所有启用状态的段落 List<AiKnowledgeSegmentDO> segments = segmentMapper.selectListByKnowledgeIdAndStatus( knowledgeId, CommonStatusEnum.ENABLE.getStatus()); if (CollUtil.isEmpty(segments)) { return; } // 2.2 遍历所有段落,重新索引 for (AiKnowledgeSegmentDO segment : segments) { // 删除旧的向量 deleteVectorStore(vectorStore, segment); // 重新创建向量 writeVectorStore(vectorStore, segment, new Document(segment.getContent())); } log.info("[reindexKnowledgeSegmentByKnowledgeId][知识库({}) 重新索引完成,共处理 {} 个段落]", knowledgeId, segments.size()); } private void writeVectorStore(VectorStore vectorStore, AiKnowledgeSegmentDO segmentDO, Document segment) { // 1. 向量存储 // 为什么要 toString 呢?因为部分 VectorStore 实现,不支持 Long 类型,例如说 QdrantVectorStore segment.getMetadata().put(VECTOR_STORE_METADATA_KNOWLEDGE_ID, segmentDO.getKnowledgeId().toString()); segment.getMetadata().put(VECTOR_STORE_METADATA_DOCUMENT_ID, segmentDO.getDocumentId().toString()); segment.getMetadata().put(VECTOR_STORE_METADATA_SEGMENT_ID, segmentDO.getId().toString()); vectorStore.add(List.of(segment)); // 2. 更新向量 ID segmentMapper.updateById(new AiKnowledgeSegmentDO().setId(segmentDO.getId()).setVectorId(segment.getId())); } private void deleteVectorStore(VectorStore vectorStore, AiKnowledgeSegmentDO segmentDO) { // 1. 更新向量 ID if (StrUtil.isEmpty(segmentDO.getVectorId())) { return; } segmentMapper.updateById(new AiKnowledgeSegmentDO().setId(segmentDO.getId()) .setVectorId(AiKnowledgeSegmentDO.VECTOR_ID_EMPTY)); // 2. 删除向量 vectorStore.delete(List.of(segmentDO.getVectorId())); } @Override public List<AiKnowledgeSegmentSearchRespBO> searchKnowledgeSegment(AiKnowledgeSegmentSearchReqBO reqBO) { // 1. 校验 AiKnowledgeDO knowledge = knowledgeService.validateKnowledgeExists(reqBO.getKnowledgeId()); // 2. 检索 List<Document> documents = searchDocument(knowledge, reqBO); if (CollUtil.isEmpty(documents)) { return ListUtil.empty(); } // 3.1 段落召回 List<AiKnowledgeSegmentDO> segments = segmentMapper .selectListByVectorIds(convertList(documents, Document::getId)); if (CollUtil.isEmpty(segments)) { return ListUtil.empty(); } // 3.2 增加召回次数 segmentMapper.updateRetrievalCountIncrByIds(convertList(segments, AiKnowledgeSegmentDO::getId)); // 4. 构建结果 List<AiKnowledgeSegmentSearchRespBO> result = convertList(segments, segment -> { Document document = CollUtil.findOne(documents, // 找到对应的文档 doc -> Objects.equals(doc.getId(), segment.getVectorId())); if (document == null) { return null; } return BeanUtils.toBean(segment, AiKnowledgeSegmentSearchRespBO.class) .setScore(document.getScore()); }); result.sort((o1, o2) -> Double.compare(o2.getScore(), o1.getScore())); // 按照分数降序排序 return result; } /** * 基于 Embedding + Rerank Model,检索知识库中的文档 * * @param knowledge 知识库 * @param reqBO 检索请求 * @return 文档列表 */ private List<Document> searchDocument(AiKnowledgeDO knowledge, AiKnowledgeSegmentSearchReqBO reqBO) { VectorStore vectorStore = getVectorStoreById(knowledge); Integer topK = ObjUtil.defaultIfNull(reqBO.getTopK(), knowledge.getTopK()); Double similarityThreshold = ObjUtil.defaultIfNull(reqBO.getSimilarityThreshold(), knowledge.getSimilarityThreshold()); // 1. 向量检索 int searchTopK = rerankModel != null ? topK * RERANK_RETRIEVAL_FACTOR : topK; double searchSimilarityThreshold = rerankModel != null ? SIMILARITY_THRESHOLD_ACCEPT_ALL : similarityThreshold; SearchRequest.Builder searchRequestBuilder = SearchRequest.builder() .query(reqBO.getContent()) .topK(searchTopK).similarityThreshold(searchSimilarityThreshold) .filterExpression(new FilterExpressionBuilder() .eq(VECTOR_STORE_METADATA_KNOWLEDGE_ID, reqBO.getKnowledgeId().toString()).build()); List<Document> documents = vectorStore.similaritySearch(searchRequestBuilder.build()); if (CollUtil.isEmpty(documents)) { return documents; } // 2. Rerank 重排序 if (rerankModel != null) { RerankResponse rerankResponse = rerankModel.call(new RerankRequest(reqBO.getContent(), documents, DashScopeRerankOptions.builder().topN(topK).build())); documents = convertList(rerankResponse.getResults(), documentWithScore -> documentWithScore.getScore() >= similarityThreshold ? documentWithScore.getOutput() : null); } return documents; } @Override public List<AiKnowledgeSegmentDO> splitContent(String url, Integer segmentMaxTokens) { // 1. 读取 URL 内容 String content = knowledgeDocumentService.readUrl(url); // 2.1 自动检测文档类型并选择策略 AiDocumentSplitStrategyEnum strategy = detectDocumentStrategy(content, url); // 2.2 文档切片 List<Document> documentSegments = splitContentByStrategy(content, segmentMaxTokens, strategy, url); // 3. 转换为段落对象 return convertList(documentSegments, segment -> { if (StrUtil.isEmpty(segment.getText())) { return null; } return new AiKnowledgeSegmentDO() .setContent(segment.getText()) .setContentLength(segment.getText().length()) .setTokens(tokenCountEstimator.estimate(segment.getText())); }); } /** * 校验段落是否存在 * * @param id 文档编号 * @return 段落信息 */ private AiKnowledgeSegmentDO validateKnowledgeSegmentExists(Long id) { AiKnowledgeSegmentDO knowledgeSegment = segmentMapper.selectById(id); if (knowledgeSegment == null) { throw exception(KNOWLEDGE_SEGMENT_NOT_EXISTS); } return knowledgeSegment; } private VectorStore getVectorStoreById(AiKnowledgeDO knowledge) { return modelService.getOrCreateVectorStore(knowledge.getEmbeddingModelId(), VECTOR_STORE_METADATA_TYPES); } private VectorStore getVectorStoreById(Long knowledgeId) { AiKnowledgeDO knowledge = knowledgeService.validateKnowledgeExists(knowledgeId); return getVectorStoreById(knowledge); } /** * 根据策略切分内容 * * @param content 文档内容 * @param segmentMaxTokens 分段的最大 Token 数 * @param strategy 切片策略 * @param url 文档 URL(用于自动检测文件类型) * @return 切片后的文档列表 */ @SuppressWarnings("EnhancedSwitchMigration") private List<Document> splitContentByStrategy(String content, Integer segmentMaxTokens, AiDocumentSplitStrategyEnum strategy, String url) { // 自动检测策略 if (strategy == AiDocumentSplitStrategyEnum.AUTO) { strategy = detectDocumentStrategy(content, url); log.info("[splitContentByStrategy][自动检测到文档策略: {}]", strategy.getName()); } // 根据策略切分 TextSplitter textSplitter; switch (strategy) { case MARKDOWN_QA: textSplitter = new MarkdownQaSplitter(segmentMaxTokens); break; case SEMANTIC: textSplitter = new SemanticTextSplitter(segmentMaxTokens); break; case PARAGRAPH: textSplitter = new SemanticTextSplitter(segmentMaxTokens, 0); // 段落切分,无重叠 break; case TOKEN: default: textSplitter = buildTokenTextSplitter(segmentMaxTokens); break; } // 执行切分 return textSplitter.apply(Collections.singletonList(new Document(content))); } /** * 自动检测文档类型并选择切片策略 * * @param content 文档内容 * @param url 文档 URL * @return 推荐的切片策略 */ private AiDocumentSplitStrategyEnum detectDocumentStrategy(String content, String url) { if (StrUtil.isEmpty(content)) { return AiDocumentSplitStrategyEnum.TOKEN; } // 1. 检测 Markdown QA 格式 if (isMarkdownQaFormat(content, url)) { return AiDocumentSplitStrategyEnum.MARKDOWN_QA; } // 2. 检测普通 Markdown 文档 if (isMarkdownDocument(url)) { return AiDocumentSplitStrategyEnum.SEMANTIC; } // 3. 默认使用语义切分(比 Token 切分更智能) return AiDocumentSplitStrategyEnum.SEMANTIC; } /** * 检测是否为 Markdown QA 格式 * 特征:包含多个二级标题(## )且标题后紧跟答案内容 */ private boolean isMarkdownQaFormat(String content, String url) { // 文件扩展名判断 if (StrUtil.isNotEmpty(url) && !url.toLowerCase().endsWith(".md")) { return false; } // 统计二级标题数量 long h2Count = content.lines() .filter(line -> line.trim().startsWith("## ")) .count(); // 要求一:至少包含 2 个二级标题才认为是 QA 格式 if (h2Count < 2) { return false; } // 要求二:检查标题占比(QA 文档标题行数相对较多),如果二级标题占比超过 10%,认为是 QA 格式 long totalLines = content.lines().count(); double h2Ratio = (double) h2Count / totalLines; return h2Ratio > 0.1; } /** * 检测是否为 Markdown 文档 */ private boolean isMarkdownDocument(String url) { return StrUtil.endWithAnyIgnoreCase(url, ".md", ".markdown"); } /** * 构建基于 Token 的文本切片器(原有逻辑保留) */ private static TextSplitter buildTokenTextSplitter(Integer segmentMaxTokens) { return TokenTextSplitter.builder() .withChunkSize(segmentMaxTokens) .withMinChunkSizeChars(Integer.MAX_VALUE) // 忽略字符的截断 .withMinChunkLengthToEmbed(1) // 允许的最小有效分段长度 .withMaxNumChunks(Integer.MAX_VALUE) .withKeepSeparator(true) // 保留分隔符 .build(); } @Override public List<AiKnowledgeSegmentProcessRespVO> getKnowledgeSegmentProcessList(List<Long> documentIds) { if (CollUtil.isEmpty(documentIds)) { return Collections.emptyList(); } return segmentMapper.selectProcessList(documentIds); } @Override public Long createKnowledgeSegment(AiKnowledgeSegmentSaveReqVO createReqVO) { // 1.1 校验文档是否存在 AiKnowledgeDocumentDO document = knowledgeDocumentService .validateKnowledgeDocumentExists(createReqVO.getDocumentId()); // 1.2 获取知识库信息 AiKnowledgeDO knowledge = knowledgeService.validateKnowledgeExists(document.getKnowledgeId()); // 1.3 校验 token 熟练 Integer tokens = tokenCountEstimator.estimate(createReqVO.getContent()); if (tokens > document.getSegmentMaxTokens()) { throw exception(KNOWLEDGE_SEGMENT_CONTENT_TOO_LONG, tokens, document.getSegmentMaxTokens()); } // 2. 保存段落 AiKnowledgeSegmentDO segment = BeanUtils.toBean(createReqVO, AiKnowledgeSegmentDO.class) .setKnowledgeId(knowledge.getId()).setDocumentId(document.getId()) .setContentLength(createReqVO.getContent().length()).setTokens(tokens) .setVectorId(AiKnowledgeSegmentDO.VECTOR_ID_EMPTY) .setRetrievalCount(0).setStatus(CommonStatusEnum.ENABLE.getStatus()); segmentMapper.insert(segment); // 3. 向量化 writeVectorStore(getVectorStoreById(knowledge), segment, new Document(segment.getContent())); return segment.getId(); } @Override public AiKnowledgeSegmentDO getKnowledgeSegment(Long id) { public AiKnowledgeSegmentDO getSegment(Long id) { return segmentMapper.selectById(id); } @Override public List<AiKnowledgeSegmentDO> getKnowledgeSegmentList(Collection<Long> ids) { if (CollUtil.isEmpty(ids)) { return Collections.emptyList(); public AiKnowledgeSegmentDO validateSegmentExists(Long id) { AiKnowledgeSegmentDO segment = segmentMapper.selectById(id); if (segment == null) throw exception(KNOWLEDGE_SEGMENT_NOT_EXISTS); return segment; } @Override public PageResult<AiKnowledgeSegmentDO> getSegmentPage(AiKnowledgeSegmentPageReqVO pageReqVO) { return segmentMapper.selectPage(pageReqVO); } @Override @Transactional(rollbackFor = Exception.class) public Long createSegment(AiKnowledgeSegmentSaveReqVO saveReqVO) { if (saveReqVO.getKnowledgeId() == null) throw exception(KNOWLEDGE_NOT_EXISTS); if (saveReqVO.getDocumentId() == null) throw exception(KNOWLEDGE_DOCUMENT_NOT_EXISTS); AiKnowledgeDO knowledge = knowledgeService.validateKnowledge(saveReqVO.getKnowledgeId()); AiKnowledgeSegmentDO segment = new AiKnowledgeSegmentDO(); segment.setKnowledgeId(saveReqVO.getKnowledgeId()); segment.setDocumentId(saveReqVO.getDocumentId()); segment.setContent(saveReqVO.getContent()); segment.setContentLength(saveReqVO.getContent().length()); segment.setTokens(estimateTokens(saveReqVO.getContent())); segment.setStatus(saveReqVO.getStatus() != null ? saveReqVO.getStatus() : CommonStatusEnum.ENABLE.getStatus()); segment.setRetrievalCount(0); segment.setVectorId(AiKnowledgeSegmentDO.VECTOR_ID_EMPTY); segmentMapper.insert(segment); // 向量化 VectorStore vectorStore = modelService.getOrCreateVectorStore( knowledge.getEmbeddingModelId(), buildMetadataFields()); Map<String, Object> metadata = new HashMap<>(); metadata.put(METADATA_KNOWLEDGE_ID, knowledge.getId()); metadata.put(METADATA_DOCUMENT_ID, segment.getDocumentId()); metadata.put(METADATA_SEGMENT_ID, segment.getId()); Document doc = new Document(segment.getId().toString(), segment.getContent(), metadata); try { vectorStore.add(Collections.singletonList(doc)); segment.setVectorId(segment.getId().toString()); segmentMapper.updateById(segment); } catch (Exception e) { log.error("手动创建分段向量化失败,segmentId={}", segment.getId(), e); segment.setStatus(CommonStatusEnum.DISABLE.getStatus()); segmentMapper.updateById(segment); } return segmentMapper.selectByIds(ids); return segment.getId(); } @Override @Transactional(rollbackFor = Exception.class) public void updateSegment(AiKnowledgeSegmentSaveReqVO saveReqVO) { AiKnowledgeSegmentDO segment = validateSegmentExists(saveReqVO.getId()); AiKnowledgeDO knowledge = knowledgeService.validateKnowledge(segment.getKnowledgeId()); // 删除旧向量 if (StrUtil.isNotEmpty(segment.getVectorId()) && !AiKnowledgeSegmentDO.VECTOR_ID_EMPTY.equals(segment.getVectorId())) { VectorStore vectorStore = modelService.getOrCreateVectorStore( knowledge.getEmbeddingModelId(), buildMetadataFields()); try { vectorStore.delete(Collections.singletonList(segment.getVectorId())); } catch (Exception e) { log.warn("删除旧向量失败: {}", segment.getVectorId(), e); } } segment.setContent(saveReqVO.getContent()); segment.setContentLength(saveReqVO.getContent().length()); segment.setTokens(estimateTokens(saveReqVO.getContent())); if (saveReqVO.getStatus() != null) segment.setStatus(saveReqVO.getStatus()); segmentMapper.updateById(segment); // 重新向量化 VectorStore vectorStore = modelService.getOrCreateVectorStore( knowledge.getEmbeddingModelId(), buildMetadataFields()); Map<String, Object> metadata = new HashMap<>(); metadata.put(METADATA_KNOWLEDGE_ID, knowledge.getId()); metadata.put(METADATA_DOCUMENT_ID, segment.getDocumentId()); metadata.put(METADATA_SEGMENT_ID, segment.getId()); Document doc = new Document(segment.getId().toString(), segment.getContent(), metadata); try { vectorStore.add(Collections.singletonList(doc)); segment.setVectorId(segment.getId().toString()); segmentMapper.updateById(segment); } catch (Exception e) { log.error("更新分段向量化失败,segmentId={}", segment.getId(), e); segment.setStatus(CommonStatusEnum.DISABLE.getStatus()); segmentMapper.updateById(segment); } } @Override @Transactional(rollbackFor = Exception.class) public void deleteSegment(Long id) { AiKnowledgeSegmentDO segment = validateSegmentExists(id); AiKnowledgeDO knowledge = knowledgeService.validateKnowledge(segment.getKnowledgeId()); if (StrUtil.isNotEmpty(segment.getVectorId()) && !AiKnowledgeSegmentDO.VECTOR_ID_EMPTY.equals(segment.getVectorId())) { VectorStore vectorStore = modelService.getOrCreateVectorStore( knowledge.getEmbeddingModelId(), buildMetadataFields()); try { vectorStore.delete(Collections.singletonList(segment.getVectorId())); } catch (Exception e) { log.warn("删除向量失败: {}", segment.getVectorId(), e); } } segmentMapper.deleteById(id); } @Override public void updateSegmentStatus(AiKnowledgeSegmentUpdateStatusReqVO updateStatusReqVO) { AiKnowledgeSegmentDO segment = validateSegmentExists(updateStatusReqVO.getId()); segment.setStatus(updateStatusReqVO.getStatus()); segmentMapper.updateById(segment); } @Override @Transactional(rollbackFor = Exception.class) public void deleteSegmentsByDocumentId(Long documentId) { List<AiKnowledgeSegmentDO> segments = segmentMapper.selectListByDocumentId(documentId); if (CollUtil.isEmpty(segments)) return; AiKnowledgeDO knowledge = knowledgeService.getKnowledge(segments.get(0).getKnowledgeId()); VectorStore vectorStore = modelService.getOrCreateVectorStore( knowledge.getEmbeddingModelId(), buildMetadataFields()); for (AiKnowledgeSegmentDO segment : segments) { if (StrUtil.isNotEmpty(segment.getVectorId()) && !AiKnowledgeSegmentDO.VECTOR_ID_EMPTY.equals(segment.getVectorId())) { try { vectorStore.delete(Collections.singletonList(segment.getVectorId())); } catch (Exception e) { log.warn("删除向量[{}]失败: {}", segment.getVectorId(), e.getMessage()); } } } segmentMapper.delete(new com.baomidou.mybatisplus.core.conditions.query.LambdaQueryWrapper<AiKnowledgeSegmentDO>() .eq(AiKnowledgeSegmentDO::getDocumentId, documentId)); } @Override @Transactional(rollbackFor = Exception.class) public void saveSegments(List<AiKnowledgeSegmentDO> segments, Long knowledgeId, Long embeddingModelId) { if (CollUtil.isEmpty(segments)) return; AiKnowledgeDO knowledge = knowledgeService.validateKnowledge(knowledgeId); VectorStore vectorStore = modelService.getOrCreateVectorStore(embeddingModelId, buildMetadataFields()); // 批量插入分段记录 for (AiKnowledgeSegmentDO segment : segments) { segmentMapper.insert(segment); } // 转换为 Spring AI Document 并写入向量库 List<Document> documents = new ArrayList<>(); for (AiKnowledgeSegmentDO segment : segments) { Map<String, Object> metadata = new HashMap<>(); metadata.put(METADATA_KNOWLEDGE_ID, knowledgeId); metadata.put(METADATA_DOCUMENT_ID, segment.getDocumentId()); metadata.put(METADATA_SEGMENT_ID, segment.getId()); Document doc = new Document(segment.getId().toString(), segment.getContent(), metadata); documents.add(doc); } try { vectorStore.add(documents); // 更新 vectorId(Milvus 返回的 ID 就是传入的 docId) for (AiKnowledgeSegmentDO segment : segments) { segment.setVectorId(segment.getId().toString()); segmentMapper.updateById(segment); } log.info("批量写入向量库成功,知识库: {}, 分段数: {}", knowledgeId, segments.size()); } catch (Exception e) { log.error("写入向量库失败,知识库: {}", knowledgeId, e); for (AiKnowledgeSegmentDO segment : segments) { segment.setStatus(CommonStatusEnum.DISABLE.getStatus()); segmentMapper.updateById(segment); } throw new RuntimeException("向量库写入失败: " + e.getMessage(), e); } } @Override public List<AiKnowledgeSegmentSearchRespBO> searchSegments(AiKnowledgeSegmentSearchReqBO searchReqBO) { AiKnowledgeDO knowledge = knowledgeService.validateKnowledge(searchReqBO.getKnowledgeId()); VectorStore vectorStore = modelService.getOrCreateVectorStore( knowledge.getEmbeddingModelId(), buildMetadataFields()); int topK = searchReqBO.getTopK() != null ? searchReqBO.getTopK() : knowledge.getTopK(); double similarityThreshold = searchReqBO.getSimilarityThreshold() != null ? searchReqBO.getSimilarityThreshold() : knowledge.getSimilarityThreshold(); Filter.Expression filterExpression = new FilterExpressionBuilder() .eq(METADATA_KNOWLEDGE_ID, searchReqBO.getKnowledgeId()) .build(); SearchRequest request = SearchRequest.builder() .query(searchReqBO.getContent()) .topK(topK) .similarityThreshold(similarityThreshold) .filterExpression(filterExpression) .build(); log.info("向量检索: query={}, topK={}, threshold={}, filter=kid={}", searchReqBO.getContent(), topK, similarityThreshold, searchReqBO.getKnowledgeId()); List<Document> results = vectorStore.similaritySearch(request); log.info("向量检索结果数: {}", results != null ? results.size() : 0); if (CollUtil.isEmpty(results)) return Collections.emptyList(); // 更新检索次数 List<String> vectorIds = results.stream() .map(Document::getId) .filter(StrUtil::isNotEmpty) .collect(Collectors.toList()); if (CollUtil.isNotEmpty(vectorIds)) { List<AiKnowledgeSegmentDO> hitSegments = segmentMapper.selectListByVectorIds(vectorIds); if (CollUtil.isNotEmpty(hitSegments)) { List<Long> hitIds = hitSegments.stream().map(AiKnowledgeSegmentDO::getId).collect(Collectors.toList()); segmentMapper.updateRetrievalCountIncrByIds(hitIds); Set<Long> docIds = hitSegments.stream().map(AiKnowledgeSegmentDO::getDocumentId).collect(Collectors.toSet()); documentMapper.updateRetrievalCountIncr(docIds); } } // 收集文档 ID 并批量查询文档名称 Set<Long> resultDocIds = results.stream() .map(doc -> convertToLong(doc.getMetadata() != null ? doc.getMetadata().get(METADATA_DOCUMENT_ID) : null)) .filter(id -> id != null) .collect(Collectors.toSet()); final Map<Long, String> docNameMap; if (CollUtil.isNotEmpty(resultDocIds)) { docNameMap = documentMapper.selectBatchIds(resultDocIds).stream() .collect(Collectors.toMap( cn.iocoder.yudao.module.ai.dal.dataobject.knowledge.AiKnowledgeDocumentDO::getId, cn.iocoder.yudao.module.ai.dal.dataobject.knowledge.AiKnowledgeDocumentDO::getName, (a, b) -> a)); } else { docNameMap = Collections.emptyMap(); } return results.stream() .map(doc -> { AiKnowledgeSegmentSearchRespBO bo = new AiKnowledgeSegmentSearchRespBO(); bo.setContent(doc.getText()); bo.setScore(doc.getScore() != null ? doc.getScore() : 0.0); if (doc.getMetadata() != null) { bo.setId(convertToLong(doc.getMetadata().get(METADATA_SEGMENT_ID))); bo.setDocumentId(convertToLong(doc.getMetadata().get(METADATA_DOCUMENT_ID))); bo.setDocumentName(docNameMap.getOrDefault(bo.getDocumentId(), "未知文档")); bo.setKnowledgeId(convertToLong(doc.getMetadata().get(METADATA_KNOWLEDGE_ID))); log.debug("检索结果元数据: id={}, docId={}, kid={}, metadataKeys={}", bo.getId(), bo.getDocumentId(), bo.getKnowledgeId(), doc.getMetadata().keySet()); } if (bo.getContent() != null) { bo.setContentLength(bo.getContent().length()); bo.setTokens(estimateTokens(bo.getContent())); } return bo; }) .collect(Collectors.toList()); } private Map<String, Class<?>> buildMetadataFields() { return MapUtil.<String, Class<?>>builder() .put(METADATA_KNOWLEDGE_ID, Long.class) .put(METADATA_DOCUMENT_ID, Long.class) .put(METADATA_SEGMENT_ID, Long.class) .build(); } /** * 将 Milvus 返回的元数据值转为 Long。 * Gson 反序列化 JSON 数字默认为 Double,需要兼容多种类型。 */ private Long convertToLong(Object value) { if (value == null) return null; if (value instanceof Long v) return v; if (value instanceof Integer v) return v.longValue(); if (value instanceof Double v) return v.longValue(); if (value instanceof Float v) return v.longValue(); if (value instanceof Number v) return v.longValue(); if (value instanceof String v) { try { return Long.valueOf(v); } catch (NumberFormatException e) { return null; } } log.warn("无法转换元数据值类型: {} = {}", value.getClass().getName(), value); return null; } private Integer estimateTokens(String text) { if (StrUtil.isEmpty(text)) return 0; int chineseChars = 0, englishWords = 0; for (char c : text.toCharArray()) { if (c >= 0x4E00 && c <= 0x9FA5) chineseChars++; } for (String word : text.split("\\s+")) { if (word.matches(".*[a-zA-Z].*")) englishWords++; } return chineseChars + (int) (englishWords * 1.3); } }