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<ChatModel>) () -> {
|
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<EmbeddingModel>) () -> {
|
if (platform == AiPlatformEnum.TONG_YI) {
|
return AiAutoConfiguration.buildTongYiEmbeddingModel(apiKey, model);
|
}
|
throw new IllegalArgumentException(StrUtil.format("不支持的平台({})", platform));
|
});
|
}
|
|
@Override
|
public VectorStore getOrCreateVectorStore(Class<? extends VectorStore> type,
|
EmbeddingModel embeddingModel,
|
Map<String, Class<?>> metadataFields) {
|
// metadataFields 参与缓存 key,确保不同知识库使用不同配置时不会复用
|
String cacheKey = buildCacheKey(VectorStore.class, embeddingModel, type, metadataFields.hashCode());
|
return Singleton.get(cacheKey, (Func0<VectorStore>) () -> {
|
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, "_"));
|
}
|
|
}
|