Spring AI 多模型供应商无缝切换实现

Spring AI 通过抽象层设计实现了多模型供应商的无缝切换,核心思想是统一接口 + 策略模式 + 配置驱动

核心架构设计

1. 统一抽象层

// 统一的 ChatModel 接口
public interface ChatModel {
    ChatResponse call(Prompt prompt);
    Stream<ChatResponse> stream(Prompt prompt);
}

// 统一的 Prompt 抽象
public class Prompt {
    private final List<Message> messages;
    private final ChatOptions options;
    // ...
}

// 统一的 Message 抽象
public interface Message {
    String getContent();
    MessageType getType();
}

2. 多供应商实现

// OpenAI 实现
public class OpenAiChatModel implements ChatModel {
    private final OpenAiApi openAiApi;
    private final OpenAiChatOptions options;
    
    @Override
    public ChatResponse call(Prompt prompt) {
        // 转换为 OpenAI 特定格式
        OpenAiApi.ChatRequest request = convertToOpenAiRequest(prompt);
        OpenAiApi.ChatResponse response = openAiApi.chatCompletion(request);
        return convertToChatResponse(response);
    }
}

// Azure OpenAI 实现
public class AzureOpenAiChatModel implements ChatModel {
    private final AzureOpenAiApi azureApi;
    
    @Override
    public ChatResponse call(Prompt prompt) {
        // 转换为 Azure OpenAI 格式
        AzureOpenAiApi.ChatRequest request = convertToAzureRequest(prompt);
        AzureOpenAiApi.ChatResponse response = azureApi.chatCompletion(request);
        return convertToChatResponse(response);
    }
}

// Anthropic Claude 实现
public class AnthropicChatModel implements ChatModel {
    private final AnthropicApi anthropicApi;
    
    @Override
    public ChatResponse call(Prompt prompt) {
        // 转换为 Claude 格式
        AnthropicApi.MessageRequest request = convertToAnthropicRequest(prompt);
        AnthropicApi.MessageResponse response = anthropicApi.message(request);
        return convertToChatResponse(response);
    }
}

配置驱动的切换机制

1. 自动配置类

@Configuration
@ConditionalOnClass(ChatModel.class)
public class ChatModelAutoConfiguration {
    
    @Bean
    @ConditionalOnMissingBean
    @ConditionalOnProperty(name = "spring.ai.chat.provider", havingValue = "openai")
    public ChatModel openAiChatModel(OpenAiProperties properties) {
        return new OpenAiChatModel(
            new OpenAiApi(properties.getApiKey(), properties.getBaseUrl()),
            properties.getOptions()
        );
    }
    
    @Bean
    @ConditionalOnMissingBean
    @ConditionalOnProperty(name = "spring.ai.chat.provider", havingValue = "azure")
    public ChatModel azureOpenAiChatModel(AzureOpenAiProperties properties) {
        return new AzureOpenAiChatModel(
            new AzureOpenAiApi(properties.getApiKey(), properties.getEndpoint()),
            properties.getOptions()
        );
    }
    
    @Bean
    @ConditionalOnMissingBean
    @ConditionalOnProperty(name = "spring.ai.chat.provider", havingValue = "anthropic")
    public ChatModel anthropicChatModel(AnthropicProperties properties) {
        return new AnthropicChatModel(
            new AnthropicApi(properties.getApiKey()),
            properties.getOptions()
        );
    }
}

2. 配置文件切换

# application.yml - 切换到 OpenAI
spring:
  ai:
    chat:
      provider: openai
    openai:
      api-key: ${OPENAI_API_KEY}
      base-url: https://api.openai.com/v1
      chat:
        options:
          model: gpt-4
          temperature: 0.7
          max-tokens: 1000

---

# application-azure.yml - 切换到 Azure OpenAI
spring:
  ai:
    chat:
      provider: azure
    azure:
      openai:
        api-key: ${AZURE_OPENAI_API_KEY}
        endpoint: https://your-resource.openai.azure.com
        deployment-name: gpt-4-deployment
        chat:
          options:
            temperature: 0.7
            max-tokens: 1000

---

# application-anthropic.yml - 切换到 Claude
spring:
  ai:
    chat:
      provider: anthropic
    anthropic:
      api-key: ${ANTHROPIC_API_KEY}
      chat:
        options:
          model: claude-3-opus-20240229
          temperature: 0.7
          max-tokens: 1000

运行时动态切换

1. 多模型管理器

@Component
public class ChatModelManager {
    
    private final Map<String, ChatModel> models = new ConcurrentHashMap<>();
    private String currentProvider;
    
    @Autowired
    public ChatModelManager(
            @Qualifier("openAiChatModel") ChatModel openAiModel,
            @Qualifier("azureOpenAiChatModel") ChatModel azureModel,
            @Qualifier("anthropicChatModel") ChatModel anthropicModel) {
        
        models.put("openai", openAiModel);
        models.put("azure", azureModel);
        models.put("anthropic", anthropicModel);
        this.currentProvider = "openai"; // 默认
    }
    
    public ChatResponse chat(Prompt prompt) {
        return getCurrentModel().call(prompt);
    }
    
    public ChatResponse chat(String provider, Prompt prompt) {
        return getModel(provider).call(prompt);
    }
    
    public void switchProvider(String provider) {
        if (!models.containsKey(provider)) {
            throw new IllegalArgumentException("Unknown provider: " + provider);
        }
        this.currentProvider = provider;
    }
    
    private ChatModel getCurrentModel() {
        return models.get(currentProvider);
    }
    
    private ChatModel getModel(String provider) {
        ChatModel model = models.get(provider);
        if (model == null) {
            throw new IllegalArgumentException("Unknown provider: " + provider);
        }
        return model;
    }
}

2. 动态切换服务

@Service
public class ChatService {
    
    @Autowired
    private ChatModelManager modelManager;
    
    // 使用默认模型
    public String chat(String message) {
        Prompt prompt = new Prompt(new UserMessage(message));
        ChatResponse response = modelManager.chat(prompt);
        return response.getResult().getOutput().getContent();
    }
    
    // 指定供应商
    public String chat(String provider, String message) {
        Prompt prompt = new Prompt(new UserMessage(message));
        ChatResponse response = modelManager.chat(provider, prompt);
        return response.getResult().getOutput().getContent();
    }
    
    // 切换供应商
    public void switchProvider(String provider) {
        modelManager.switchProvider(provider);
    }
    
    // 根据场景自动选择
    public String smartChat(String message, ChatScenario scenario) {
        String provider = selectProvider(scenario);
        return chat(provider, message);
    }
    
    private String selectProvider(ChatScenario scenario) {
        return switch (scenario) {
            case CODE_GENERATION -> "openai";  // GPT-4 擅长代码
            case LONG_CONTEXT -> "anthropic";  // Claude 上下文更长
            case ENTERPRISE -> "azure";        // Azure 企业级
            default -> "openai";
        };
    }
}

请求转换适配器

1. 统一请求转换

public interface RequestAdapter<T> {
    T adapt(Prompt prompt, ChatOptions options);
}

// OpenAI 适配器
@Component
public class OpenAiRequestAdapter implements RequestAdapter<OpenAiApi.ChatRequest> {
    
    @Override
    public OpenAiApi.ChatRequest adapt(Prompt prompt, ChatOptions options) {
        List<OpenAiApi.ChatMessage> messages = prompt.getInstructions().stream()
            .map(this::convertMessage)
            .collect(Collectors.toList());
        
        return OpenAiApi.ChatRequest.builder()
            .messages(messages)
            .model(options.getModel())
            .temperature(options.getTemperature())
            .maxTokens(options.getMaxTokens())
            .build();
    }
    
    private OpenAiApi.ChatMessage convertMessage(Message message) {
        return new OpenAiApi.ChatMessage(
            message.getContent(),
            OpenAiApi.ChatMessageRole.valueOf(message.getMessageType().name())
        );
    }
}

// Anthropic 适配器
@Component
public class AnthropicRequestAdapter implements RequestAdapter<AnthropicApi.MessageRequest> {
    
    @Override
    public AnthropicApi.MessageRequest adapt(Prompt prompt, ChatOptions options) {
        List<AnthropicApi.Message> messages = prompt.getInstructions().stream()
            .map(this::convertMessage)
            .collect(Collectors.toList());
        
        return AnthropicApi.MessageRequest.builder()
            .messages(messages)
            .model(options.getModel())
            .temperature(options.getTemperature())
            .maxTokens(options.getMaxTokens())
            .build();
    }
    
    private AnthropicApi.Message convertMessage(Message message) {
        return AnthropicApi.Message.builder()
            .role(message.getMessageType().name().toLowerCase())
            .content(message.getContent())
            .build();
    }
}

2. 响应转换适配器

public interface ResponseAdapter<T> {
    ChatResponse adapt(T response);
}

@Component
public class OpenAiResponseAdapter implements ResponseAdapter<OpenAiApi.ChatResponse> {
    
    @Override
    public ChatResponse adapt(OpenAiApi.ChatResponse response) {
        List<Generation> generations = response.getChoices().stream()
            .map(choice -> new Generation(
                new AssistantMessage(choice.getMessage().getContent()),
                choice.getFinishReason()
            ))
            .collect(Collectors.toList());
        
        return new ChatResponse(generations, response.getUsage());
    }
}

@Component
public class AnthropicResponseAdapter implements ResponseAdapter<AnthropicApi.MessageResponse> {
    
    @Override
    public ChatResponse adapt(AnthropicApi.MessageResponse response) {
        List<Generation> generations = List.of(
            new Generation(
                new AssistantMessage(response.getContent()[0].getText()),
                response.getStopReason()
            )
        );
        
        return new ChatResponse(generations, response.getUsage());
    }
}

统一选项抽象

1. ChatOptions 层次结构

// 基础选项
public class ChatOptions {
    private String model;
    private Double temperature;
    private Integer maxTokens;
    private Double topP;
    private List<String> stop;
    // getters and setters
}

// OpenAI 特定选项
public class OpenAiChatOptions extends ChatOptions {
    private String frequencyPenalty;
    private String presencePenalty;
    private String responseFormat;
    // OpenAI 特有属性
}

// Anthropic 特定选项
public class AnthropicChatOptions extends ChatOptions {
    private Integer topK;
    private List<String> stopSequences;
    private String anthropicVersion;
    // Anthropic 特有属性
}

2. 选项验证和转换

@Component
public class ChatOptionsValidator {
    
    public void validate(ChatOptions options, String provider) {
        switch (provider) {
            case "openai" -> validateOpenAiOptions(options);
            case "azure" -> validateAzureOptions(options);
            case "anthropic" -> validateAnthropicOptions(options);
        }
    }
    
    private void validateOpenAiOptions(ChatOptions options) {
        if (options.getTemperature() != null && 
            (options.getTemperature() < 0 || options.getTemperature() > 2)) {
            throw new IllegalArgumentException(
                "OpenAI temperature must be between 0 and 2"
            );
        }
    }
    
    private void validateAnthropicOptions(ChatOptions options) {
        if (options.getTemperature() != null && 
            (options.getTemperature() < 0 || options.getTemperature() > 1)) {
            throw new IllegalArgumentException(
                "Anthropic temperature must be between 0 and 1"
            );
        }
    }
}

完整使用示例

1. Controller 层

@RestController
@RequestMapping("/api/chat")
public class ChatController {
    
    @Autowired
    private ChatService chatService;
    
    // 使用默认模型
    @PostMapping
    public ResponseEntity<ChatResponse> chat(@RequestBody ChatRequest request) {
        String response = chatService.chat(request.getMessage());
        return ResponseEntity.ok(new ChatResponse(response));
    }
    
    // 指定供应商
    @PostMapping("/{provider}")
    public ResponseEntity<ChatResponse> chat(
            @PathVariable String provider,
            @RequestBody ChatRequest request) {
        String response = chatService.chat(provider, request.getMessage());
        return ResponseEntity.ok(new ChatResponse(response));
    }
    
    // 切换供应商
    @PostMapping("/switch/{provider}")
    public ResponseEntity<Void> switchProvider(@PathVariable String provider) {
        chatService.switchProvider(provider);
        return ResponseEntity.ok().build();
    }
    
    // 智能选择
    @PostMapping("/smart")
    public ResponseEntity<ChatResponse> smartChat(
            @RequestBody SmartChatRequest request) {
        String response = chatService.smartChat(
            request.getMessage(), 
            request.getScenario()
        );
        return ResponseEntity.ok(new ChatResponse(response));
    }
}

2. 配置类

@Configuration
@EnableConfigurationProperties(AiProperties.class)
public class AiConfiguration {
    
    @Bean
    @ConditionalOnProperty(name = "spring.ai.multi-model.enabled", havingValue = "true")
    public ChatModelManager multiModelManager(
            ApplicationContext context) {
        
        Map<String, ChatModel> models = context.getBeansOfType(ChatModel.class);
        ChatModelManager manager = new ChatModelManager();
        
        models.forEach((name, model) -> {
            String provider = extractProviderName(name);
            manager.registerModel(provider, model);
        });
        
        return manager;
    }
    
    private String extractProviderName(String beanName) {
        if (beanName.contains("OpenAi")) return "openai";
        if (beanName.contains("Azure")) return "azure";
        if (beanName.contains("Anthropic")) return "anthropic";
        return "default";
    }
}

@ConfigurationProperties(prefix = "spring.ai")
public class AiProperties {
    private String chatProvider;
    private MultiModel multiModel = new MultiModel();
    private OpenAiProperties openai = new OpenAiProperties();
    private AzureOpenAiProperties azure = new AzureOpenAiProperties();
    private AnthropicProperties anthropic = new AnthropicProperties();
    
    // getters and setters
    
    public static class MultiModel {
        private boolean enabled = false;
        private String defaultProvider = "openai";
        // getters and setters
    }
}

测试示例

@SpringBootTest
class ChatServiceTest {
    
    @Autowired
    private ChatService chatService;
    
    @Test
    void testOpenAiChat() {
        String response = chatService.chat("openai", "Hello, OpenAI!");
        assertNotNull(response);
        assertFalse(response.isEmpty());
    }
    
    @Test
    void testAnthropicChat() {
        String response = chatService.chat("anthropic", "Hello, Claude!");
        assertNotNull(response);
        assertFalse(response.isEmpty());
    }
    
    @Test
    void testProviderSwitch() {
        chatService.switchProvider("openai");
        String response1 = chatService.chat("Test message");
        
        chatService.switchProvider("anthropic");
        String response2 = chatService.chat("Test message");
        
        // 两个响应可能不同,因为使用了不同的模型
        assertNotNull(response1);
        assertNotNull(response2);
    }
    
    @Test
    void testSmartSelection() {
        String codeResponse = chatService.smartChat(
            "Write a Python function", 
            ChatScenario.CODE_GENERATION
        );
        
        String longContextResponse = chatService.smartChat(
            "Analyze this long text...", 
            ChatScenario.LONG_CONTEXT
        );
        
        assertNotNull(codeResponse);
        assertNotNull(longContextResponse);
    }
}

总结

Spring AI 实现多模型供应商无缝切换的核心机制:

  1. 统一抽象层:定义统一的 ChatModelPromptMessage 接口
  2. 策略模式:每个供应商实现相同的接口
  3. 配置驱动:通过配置文件选择供应商
  4. 适配器模式:转换不同供应商的请求/响应格式
  5. 动态切换:运行时可以切换供应商
  6. 选项抽象:统一不同供应商的配置选项

这种设计使得开发者可以:

  • 通过修改配置文件轻松切换供应商
  • 在运行时动态选择不同的模型
  • 为不同场景使用最适合的模型
  • 最小化代码改动,提高可维护性
Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐