Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,7 @@
import org.springframework.ai.retry.RetryUtils;
import org.springframework.http.client.HttpComponentsClientHttpRequestFactory;
import org.springframework.http.client.reactive.ReactorClientHttpConnector;
import org.springframework.retry.support.RetryTemplate;
import org.springframework.stereotype.Service;
import org.springframework.util.Assert;
import org.springframework.util.StringUtils;
Expand All @@ -48,11 +49,16 @@
@RequiredArgsConstructor
public class DynamicModelFactory {

public static final RetryTemplate NO_RETRY_TEMPLATE = RetryTemplate.builder().maxAttempts(1).build();

/**
* 统一使用 OpenAiChatModel,通过 baseUrl 实现多厂商兼容
*/
public ChatModel createChatModel(ModelConfigDTO config) {
return createChatModel(config, RetryUtils.DEFAULT_RETRY_TEMPLATE);
}

public ChatModel createChatModel(ModelConfigDTO config, RetryTemplate retryTemplate) {
log.info("Creating NEW ChatModel instance. Provider: {}, Model: {}, BaseUrl: {}", config.getProvider(),
config.getModelName(), config.getBaseUrl());
// 1. 验证参数
Expand All @@ -79,13 +85,21 @@ public ChatModel createChatModel(ModelConfigDTO config) {
.streamUsage(true)
.build();
// 4. 返回统一的 OpenAiChatModel
return OpenAiChatModel.builder().openAiApi(openAiApi).defaultOptions(openAiChatOptions).build();
return OpenAiChatModel.builder()
.openAiApi(openAiApi)
.defaultOptions(openAiChatOptions)
.retryTemplate(retryTemplate)
.build();
}

/**
* Embedding 同理
*/
public EmbeddingModel createEmbeddingModel(ModelConfigDTO config) {
return createEmbeddingModel(config, RetryUtils.DEFAULT_RETRY_TEMPLATE);
}

public EmbeddingModel createEmbeddingModel(ModelConfigDTO config, RetryTemplate retryTemplate) {
log.info("Creating NEW EmbeddingModel instance. Provider: {}, Model: {}, BaseUrl: {}", config.getProvider(),
config.getModelName(), config.getBaseUrl());
checkBasic(config);
Expand All @@ -103,8 +117,7 @@ public EmbeddingModel createEmbeddingModel(ModelConfigDTO config) {

OpenAiApi openAiApi = apiBuilder.build();
return new OpenAiEmbeddingModel(openAiApi, MetadataMode.EMBED,
OpenAiEmbeddingOptions.builder().model(config.getModelName()).build(),
RetryUtils.DEFAULT_RETRY_TEMPLATE);
OpenAiEmbeddingOptions.builder().model(config.getModelName()).build(), retryTemplate);
}

private static void checkBasic(ModelConfigDTO config) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -135,7 +135,7 @@ private void testChatModel(ModelConfigDTO config) {
config.getModelName());

// 1. 创建临时模型
ChatModel tempModel = modelFactory.createChatModel(config);
ChatModel tempModel = modelFactory.createChatModel(config, DynamicModelFactory.NO_RETRY_TEMPLATE);

// 2. 发起最轻量的请求
String promptText = "Hello";
Expand All @@ -154,7 +154,7 @@ private void testEmbeddingModel(ModelConfigDTO config) {
log.info("Testing Embedding Model connection, provider: {} modelName: {}", config.getProvider(),
config.getModelName());
// 1. 创建临时模型
EmbeddingModel tempModel = modelFactory.createEmbeddingModel(config);
EmbeddingModel tempModel = modelFactory.createEmbeddingModel(config, DynamicModelFactory.NO_RETRY_TEMPLATE);

// 2. 发起请求
float[] embedding = tempModel.embed("Test");
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -114,7 +114,7 @@ void testTestConnection_chat() {
dto.setModelName("gpt-4");

ChatModel chatModel = mock(ChatModel.class);
when(modelFactory.createChatModel(dto)).thenReturn(chatModel);
when(modelFactory.createChatModel(dto, DynamicModelFactory.NO_RETRY_TEMPLATE)).thenReturn(chatModel);
when(chatModel.call("Hello")).thenReturn("Hi there");

assertDoesNotThrow(() -> service.testConnection(dto));
Expand All @@ -128,7 +128,7 @@ void testTestConnection_embedding() {
dto.setModelName("text-embedding");

EmbeddingModel embeddingModel = mock(EmbeddingModel.class);
when(modelFactory.createEmbeddingModel(dto)).thenReturn(embeddingModel);
when(modelFactory.createEmbeddingModel(dto, DynamicModelFactory.NO_RETRY_TEMPLATE)).thenReturn(embeddingModel);
when(embeddingModel.embed("Test")).thenReturn(new float[] { 0.1f, 0.2f });

assertDoesNotThrow(() -> service.testConnection(dto));
Expand All @@ -148,7 +148,7 @@ void testTestConnection_chatReturnsEmpty() {
dto.setModelType("CHAT");

ChatModel chatModel = mock(ChatModel.class);
when(modelFactory.createChatModel(dto)).thenReturn(chatModel);
when(modelFactory.createChatModel(dto, DynamicModelFactory.NO_RETRY_TEMPLATE)).thenReturn(chatModel);
when(chatModel.call("Hello")).thenReturn("");

assertThrows(RuntimeException.class, () -> service.testConnection(dto));
Expand All @@ -160,7 +160,7 @@ void testTestConnection_embeddingReturnsEmpty() {
dto.setModelType("EMBEDDING");

EmbeddingModel embeddingModel = mock(EmbeddingModel.class);
when(modelFactory.createEmbeddingModel(dto)).thenReturn(embeddingModel);
when(modelFactory.createEmbeddingModel(dto, DynamicModelFactory.NO_RETRY_TEMPLATE)).thenReturn(embeddingModel);
when(embeddingModel.embed("Test")).thenReturn(new float[0]);

assertThrows(RuntimeException.class, () -> service.testConnection(dto));
Expand Down
Loading