package cn.iocoder.yudao.module.ai.framework.ai.core.model; import cn.hutool.core.lang.Singleton; import cn.hutool.core.lang.func.Func0; import cn.hutool.core.util.ArrayUtil; 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 lombok.extern.slf4j.Slf4j; import org.springframework.ai.chat.model.ChatModel; import org.springframework.ai.embedding.BatchingStrategy; import org.springframework.ai.embedding.EmbeddingModel; 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.MilvusServiceClientProperties; import org.springframework.ai.vectorstore.milvus.autoconfigure.MilvusVectorStoreProperties; import java.util.Map; /** * AI Model 模型工厂实现类 * * 使用 OpenAI 兼容接口对接通义千问 + Milvus 向量存储 */ @Slf4j public class AiModelFactoryImpl implements AiModelFactory { @Override 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) () -> { if (platform == AiPlatformEnum.TONG_YI) { return AiAutoConfiguration.buildTongYiChatModel(apiKey, model); } throw new IllegalArgumentException(StrUtil.format("不支持的平台({})", platform)); }); } @Override public ChatModel getDefaultChatModel(AiPlatformEnum platform) { if (platform == AiPlatformEnum.TONG_YI) { return SpringUtil.getBean(OpenAiChatModel.class); } throw new IllegalArgumentException(StrUtil.format("不支持的平台({})", platform)); } @Override 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) () -> { if (platform == AiPlatformEnum.TONG_YI) { return AiAutoConfiguration.buildTongYiEmbeddingModel(apiKey, model); } throw new IllegalArgumentException(StrUtil.format("不支持的平台({})", platform)); }); } @Override public VectorStore getOrCreateVectorStore(Class type, EmbeddingModel embeddingModel, Map> metadataFields) { // metadataFields 参与缓存 key,确保不同知识库使用不同配置时不会复用 String cacheKey = buildCacheKey(VectorStore.class, embeddingModel, type, metadataFields.hashCode()); return Singleton.get(cacheKey, (Func0) () -> { if (type == MilvusVectorStore.class) { return buildMilvusVectorStore(embeddingModel); } throw new IllegalArgumentException(StrUtil.format("不支持的向量存储类型({})", type)); }); } private MilvusVectorStore buildMilvusVectorStore(EmbeddingModel embeddingModel) { MilvusVectorStoreProperties serverProperties = SpringUtil.getBean(MilvusVectorStoreProperties.class); MilvusServiceClientProperties clientProperties = SpringUtil.getBean(MilvusServiceClientProperties.class); 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); 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 String buildCacheKey(Class clazz, Object... params) { if (ArrayUtil.isEmpty(params)) return clazz.getName(); return StrUtil.format("{}#{}", clazz.getName(), ArrayUtil.join(params, "_")); } }