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 cn.iocoder.yudao.module.ai.framework.ai.config.YudaoAiProperties; import lombok.extern.slf4j.Slf4j; import org.springframework.ai.chat.model.ChatModel; 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 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) { 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) { YudaoAiProperties aiProperties = SpringUtil.getBean(YudaoAiProperties.class); YudaoAiProperties.VectorStore.Milvus milvusConfig = aiProperties.getVectorStore().getMilvus(); var connectParam = io.milvus.param.ConnectParam.newBuilder() .withHost(milvusConfig.getHost()) .withPort(milvusConfig.getPort()) .withDatabaseName(milvusConfig.getDatabaseName()) .build(); var milvusClient = new io.milvus.client.MilvusServiceClient(connectParam); MilvusVectorStore vectorStore = MilvusVectorStore.builder(milvusClient, embeddingModel) .databaseName(milvusConfig.getDatabaseName()) .collectionName(milvusConfig.getCollectionName()) .initializeSchema(milvusConfig.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, "_")); } }