Spring AI多模型路由与动态切换

Spring AI多模型路由与动态切换

前置知识

  • Spring AI基础配置
  • 了解不同LLM模型的特点与价格
  • Spring AOP与Route界定
  • 负载均衡基本概念

核心概念

在生产环境中,单一模型往往无法满足所有场景需求。多模型路由允许根据请求特征(复杂度、类型、用户等级、延迟要求)选择最合适的模型,同时实现容灾、降级和成本控制。

多模型路由架构

复制代码
请求 → 路由策略评估 → 模型选择 → API调用
                                    ↓
                            ← 失败 → 降级到备用模型

完整实现

1. 多模型ChatModel注册

java 复制代码
@Configuration
public class MultiModelConfig {

    /**
     * 主力模型 - 通义千问Plus
     */
    @Bean("qwemPlusChatModel")
    @Primary
    public ChatModel qwemPlusChatModel() {
        DashScopeApi api = new DashScopeApi(
            System.getenv("DASHSCOPE_API_KEY")
        );
        
        DashScopeChatOptions defaultOptions = DashScopeChatOptions.builder()
                .withModel("qwen-plus")
                .withTemperature(0.7)
                .withMaxTokens(4096)
                .build();
        
        return new DashScopeChatModel(api, defaultOptions);
    }

    /**
     * 快速模型 - Qwen Turbo (低成本快速响应)
     */
    @Bean("qwenTurboChatModel")
    public ChatModel qwenTurboChatModel() {
        DashScopeApi api = new DashScopeApi(
            System.getenv("DASHSCOPE_API_KEY")
        );
        
        DashScopeChatOptions defaultOptions = DashScopeChatOptions.builder()
                .withModel("qwen-turbo")
                .withTemperature(0.7)
                .withMaxTokens(2048)
                .build();
        
        return new DashScopeChatModel(api, defaultOptions);
    }

    /**
     * 代码模型 - Qwen Coder
     */
    @Bean("qwenCoderChatModel")
    public ChatModel qwenCoderChatModel() {
        DashScopeApi api = new DashScopeApi(
            System.getenv("DASHSCOPE_API_KEY")
        );
        
        DashScopeChatOptions defaultOptions = DashScopeChatOptions.builder()
                .withModel("qwen-coder")
                .withTemperature(0.2)
                .withMaxTokens(4096)
                .build();
        
        return new DashScopeChatModel(api, defaultOptions);
    }

    /**
     * 备用模型 - DeepSeek
     */
    @Bean("deepSeekChatModel")
    public ChatModel deepSeekChatModel() {
        OpenAiApi api = new OpenAiApi(
            "https://api.deepseek.com",
            System.getenv("DEEPSEEK_API_KEY")
        );
        
        OpenAiChatOptions defaultOptions = OpenAiChatOptions.builder()
                .withModel("deepseek-chat")
                .withTemperature(0.7)
                .withMaxTokens(4096)
                .build();
        
        return new OpenAiChatModel(api, defaultOptions);
    }

    /**
     * 模型注册表
     */
    @Bean
    public ModelRouter modelRouter(
            @Qualifier("qwemPlusChatModel") ChatModel qwemPlus,
            @Qualifier("qwenTurboChatModel") ChatModel qwenTurbo,
            @Qualifier("qwenCoderChatModel") ChatModel qwenCoder,
            @Qualifier("deepSeekChatModel") ChatModel deepSeek) {
        
        Map<String, ChatModel> models = Map.of(
            "qwen-plus", qwemPlus,
            "qwen-turbo", qwenTurbo,
            "qwen-coder", qwenCoder,
            "deepseek-chat", deepSeek
        );
        
        return new ModelRouter(models);
    }
}

2. 模型路由核心逻辑

java 复制代码
@Slf4j
@Component
public class ModelRouter {

    private final Map<String, ChatModel> models;
    private final AtomicReference<String> currentPrimaryModel;
    private final Map<String, ModelMetrics> metrics = new ConcurrentHashMap<>();

    public ModelRouter(Map<String, ChatModel> models) {
        this.models = models;
        this.currentPrimaryModel = new AtomicReference<>("qwen-plus");
        
        // 初始化指标
        models.keySet().forEach(name -> 
            metrics.put(name, new ModelMetrics()));
    }

    /**
     * 路由策略: 根据请求内容选择模型
     */
    public ChatModel route(RouteContext context) {
        // 1. 用户指定模型
        if (context.preferredModel() != null 
            && models.containsKey(context.preferredModel())) {
            String preferred = context.preferredModel();
            if (!metrics.get(preferred).isCircuitOpen()) {
                return models.get(preferred);
            }
            log.warn("指定模型 {} 熔断中,尝试降级", preferred);
        }

        // 2. 根据内容类型路由
        if (context.isCodeRequest()) {
            return routeWithFallback("qwen-coder", "qwen-plus");
        }
        
        if (context.isCreativeRequest()) {
            return routeWithFallback("qwen-plus", "deepseek-chat");
        }
        
        if (context.isFactualRequest()) {
            return routeWithFallback("qwen-plus", "deepseek-chat");
        }

        // 3. 根据复杂度路由
        if (context.complexity() < 3) {
            return routeWithFallback("qwen-turbo", "qwen-plus");
        }

        // 4. 默认主力模型
        return routeWithFallback(currentPrimaryModel.get(), "deepseek-chat");
    }

    /**
     * 带故障转移的路由
     */
    private ChatModel routeWithFallback(String primary, String fallback) {
        ModelMetrics primaryMetrics = metrics.get(primary);
        
        if (primaryMetrics != null && primaryMetrics.isCircuitOpen()) {
            log.warn("模型 {} 已熔断,使用备用模型 {}", primary, fallback);
            
            ModelMetrics fallbackMetrics = metrics.get(fallback);
            if (fallbackMetrics != null && fallbackMetrics.isCircuitOpen()) {
                // 都熔断,选择失败率最低的
                return getLeastFailingModel();
            }
            return models.get(fallback);
        }
        
        ChatModel model = models.get(primary);
        if (model == null) {
            return models.get(fallback);
        }
        return model;
    }

    /**
     * 选择失败率最低的模型
     */
    private ChatModel getLeastFailingModel() {
        return metrics.entrySet().stream()
                .min(Comparator.comparingDouble(e -> e.getValue().getErrorRate()))
                .map(e -> models.get(e.getKey()))
                .orElse(models.values().iterator().next());
    }

    /**
     * 动态修改当前主力模型
     */
    public void switchPrimaryModel(String modelName) {
        if (models.containsKey(modelName)) {
            String oldModel = currentPrimaryModel.getAndSet(modelName);
            log.info("模型切换: {} -> {}", oldModel, modelName);
        }
    }

    /**
     * 记录成功调用
     */
    public void recordSuccess(String modelName, long durationMs) {
        ModelMetrics m = metrics.get(modelName);
        if (m != null) {
            m.recordSuccess(durationMs);
        }
    }

    /**
     * 记录失败调用
     */
    public void recordFailure(String modelName, String errorType) {
        ModelMetrics m = metrics.get(modelName);
        if (m != null) {
            m.recordFailure(errorType);
        }
    }

    /**
     * 获取模型健康状态
     */
    public Map<String, ModelHealth> getHealthStatus() {
        Map<String, ModelHealth> status = new HashMap<>();
        metrics.forEach((name, m) -> {
            status.put(name, new ModelHealth(
                name,
                m.isCircuitOpen() ? "CIRCUIT_OPEN" : m.getErrorRate() > 0.1 
                    ? "DEGRADED" : "HEALTHY",
                m.getTotalCalls(),
                m.getErrorRate(),
                m.getAvgLatencyMs()
            ));
        });
        return status;
    }
}

/**
 * 路由上下文
 */
public record RouteContext(
    String userMessage,
    String userId,
    String conversationId,
    String preferredModel,
    int complexity
) {
    public boolean isCodeRequest() {
        if (userMessage == null) return false;
        return userMessage.matches("(?i).*(代码|code|function|class|算法|编程|bug|debug|实现).*")
            || userMessage.contains("```");
    }

    public boolean isCreativeRequest() {
        if (userMessage == null) return false;
        return userMessage.matches("(?i).*(创意|故事|创作|写诗|小说|营销|广告|社交媒体|romantic).*");
    }

    public boolean isFactualRequest() {
        if (userMessage == null) return false;
        return userMessage.matches("(?i).*(是什么|解释|定义|原理|科学|研究|事实|真理).*");
    }
}

3. 模型指标与熔断器

java 复制代码
/**
 * 模型调用指标 - 滑动窗口计数器
 */
@Slf4j
public class ModelMetrics {

    private final AtomicInteger totalCalls = new AtomicInteger(0);
    private final AtomicInteger errorCalls = new AtomicInteger(0);
    private final ConcurrentLinkedDeque<Long> latentCalls = 
        new ConcurrentLinkedDeque<>();
    private final ConcurrentLinkedDeque<Boolean> recentResults = 
        new ConcurrentLinkedDeque<>();
    
    private volatile boolean circuitOpen = false;
    private volatile long circuitOpenTime = 0;
    
    private static final int WINDOW_SIZE = 100;
    private static final double CIRCUIT_OPEN_THRESHOLD = 0.5;
    private static final long CIRCUIT_OPEN_DURATION_MS = 30_000; // 30秒后尝试半开

    public void recordSuccess(long durationMs) {
        totalCalls.incrementAndGet();
        latentCalls.add(durationMs);
        recentResults.add(true);
        maintainWindow();
        
        // 恢复熔断
        if (circuitOpen && System.currentTimeMillis() > circuitOpenTime 
                + CIRCUIT_OPEN_DURATION_MS) {
            circuitOpen = false;
            log.info("模型熔断恢复,进入半开状态");
        }
    }

    public void recordFailure(String errorType) {
        totalCalls.incrementAndGet();
        errorCalls.incrementAndGet();
        recentResults.add(false);
        maintainWindow();
        
        // 检查熔断阈值
        if (getErrorRate() >= CIRCUIT_OPEN_THRESHOLD 
                && totalCalls.get() >= 10) {
            circuitOpen = true;
            circuitOpenTime = System.currentTimeMillis();
            log.warn("模型熔断开启,错误率: {:.2%}", getErrorRate());
        }
    }

    public boolean isCircuitOpen() {
        if (!circuitOpen) return false;
        
        // 熔断时间窗口过后进入半开状态,允许一个请求尝试
        if (System.currentTimeMillis() > circuitOpenTime + CIRCUIT_OPEN_DURATION_MS) {
            return false; // 半开状态
        }
        return true;
    }

    public double getErrorRate() {
        int total = totalCalls.get();
        if (total == 0) return 0.0;
        return (double) errorCalls.get() / total;
    }

    public double getAvgLatencyMs() {
        if (latentCalls.isEmpty()) return 0;
        return latentCalls.stream()
                .mapToLong(Long::longValue)
                .average()
                .orElse(0);
    }

    public int getTotalCalls() {
        return totalCalls.get();
    }

    private void maintainWindow() {
        if (recentResults.size() > WINDOW_SIZE) {
            recentResults.pollFirst();
        }
        while (latentCalls.size() > WINDOW_SIZE) {
            latentCalls.pollFirst();
        }
    }

    public double getRecentErrorRate() {
        if (recentResults.isEmpty()) return 0;
        long failCount = recentResults.stream().filter(r -> !r).count();
        return (double) failCount / recentResults.size();
    }
}

public record ModelHealth(
    String modelName,
    String status,
    int totalCalls,
    double errorRate,
    double avgLatencyMs
) {}

4. 动态模型切换的AOP实现

java 复制代码
/**
 * 路由拦截器 - 自动为ChatClient调用选择模型
 */
@Aspect
@Component
@Slf4j
public class ModelRouteAspect {

    private final ModelRouter router;
    private final MeterRegistry meterRegistry;

    public ModelRouteAspect(ModelRouter router, MeterRegistry meterRegistry) {
        this.router = router;
        this.meterRegistry = meterRegistry;
    }

    @Around("execution(* com.example.ai.service.*.*(..)) && " +
            "@annotation(routeToModel)")
    public Object routeModel(ProceedingJoinPoint pjp, 
                              RouteToModel routeToModel) throws Throwable {
        RouteContext context = extractContext(pjp.getArgs());
        ChatModel selected = router.route(context);
        String modelName = getModelName(selected);
        
        Transaction transaction = Cat.newTransaction("LLM", modelName);
        long startTime = System.currentTimeMillis();
        
        try {
            Object result = pjp.proceed();
            long duration = System.currentTimeMillis() - startTime;
            
            router.recordSuccess(modelName, duration);
            meterRegistry.timer("llm.route.success")
                    .tag("model", modelName)
                    .record(duration, TimeUnit.MILLISECONDS);
            transaction.setStatus(Transaction.SUCCESS);
            
            return result;
        } catch (Exception e) {
            long duration = System.currentTimeMillis() - startTime;
            router.recordFailure(modelName, e.getClass().getSimpleName());
            meterRegistry.counter("llm.route.failure")
                    .tag("model", modelName)
                    .tag("error", e.getClass().getSimpleName())
                    .increment();
            transaction.setStatus(e);
            throw e;
        } finally {
            transaction.complete();
        }
    }

    private RouteContext extractContext(Object[] args) {
        // 从方法参数构造路由上下文
        if (args.length > 0 && args[0] instanceof ChatRequest request) {
            return new RouteContext(
                request.message(),
                request.userId(),
                request.conversationId(),
                request.preferredModel(),
                estimateComplexity(request.message())
            );
        }
        return new RouteContext("", "", null, null, 5);
    }

    private int estimateComplexity(String message) {
        if (message == null) return 5;
        int score = 5;
        if (message.length() > 500) score += 2;
        if (message.contains("复杂") || message.contains("详细")) score += 2;
        if (message.contains("简洁") || message.contains("简短")) score -= 2;
        return Math.max(1, Math.min(10, score));
    }

    private String getModelName(ChatModel model) {
        try {
            Field field = model.getClass().getDeclaredField("chatOptions");
            field.setAccessible(true);
            ChatOptions options = (ChatOptions) field.get(model);
            return options != null ? options.getModel() : "unknown";
        } catch (Exception e) {
            return "unknown";
        }
    }
}

@Target(ElementType.METHOD)
@Retention(RetentionPolicy.RUNTIME)
public @interface RouteToModel {
}

5. 基于使用量的模型调度

java 复制代码
@Service
public class QuotaBasedScheduler {

    private final ModelRouter router;
    private final Map<String, AtomicInteger> dailyUsage = new ConcurrentHashMap<>();
    private final Map<String, Integer> dailyLimits = Map.of(
        "qwen-turbo", 10000,
        "qwen-plus", 5000,
        "qwen-coder", 3000,
        "deepseek-chat", 2000
    );

    public QuotaBasedScheduler(ModelRouter router) {
        this.router = router;
        
        // 每天重置配额
        ScheduledExecutorService scheduler = Executors.newSingleThreadScheduledExecutor();
        scheduler.scheduleAtFixedRate(
            dailyUsage::clear,
            getSecondsUntilMidnight(), 
            TimeUnit.DAYS.toSeconds(1),
            TimeUnit.SECONDS
        );
    }

    /**
     * 配额感知的模型路由
     */
    public ChatModel routeWithQuota(RouteContext context) {
        ChatModel preferred = router.route(context);
        String modelName = getModelName(preferred);
        
        AtomicInteger usage = dailyUsage.computeIfAbsent(
            modelName, k -> new AtomicInteger(0));
        int limit = dailyLimits.getOrDefault(modelName, 1000);
        
        if (usage.incrementAndGet() > limit) {
            log.warn("模型 {} 配额已用尽 ({}), 尝试降级", modelName, limit);
            // 选择还有配额的模型
            return getNextAvailableModel(modelName);
        }
        
        return preferred;
    }

    private ChatModel getNextAvailableModel(String exhaustedModel) {
        for (Map.Entry<String, Integer> entry : dailyLimits.entrySet()) {
            if (entry.getKey().equals(exhaustedModel)) continue;
            
            AtomicInteger usage = dailyUsage.computeIfAbsent(
                entry.getKey(), k -> new AtomicInteger(0));
            
            if (usage.get() < entry.getValue()) {
                return router.getModel(entry.getKey());
            }
        }
        
        // 全部配额用尽
        log.error("所有模型配额均已用尽");
        return null;
    }

    private long getSecondsUntilMidnight() {
        LocalDateTime now = LocalDateTime.now();
        LocalDateTime midnight = now.toLocalDate().plusDays(1).atStartOfDay();
        return Duration.between(now, midnight).getSeconds();
    }

    /**
     * 获取当前配额使用情况
     */
    public Map<String, QuotaStatus> getQuotaStatus() {
        Map<String, QuotaStatus> status = new HashMap<>();
        dailyLimits.forEach((model, limit) -> {
            int used = dailyUsage.getOrDefault(model, new AtomicInteger(0)).get();
            status.put(model, new QuotaStatus(model, used, limit, limit - used));
        });
        return status;
    }

    public record QuotaStatus(
        String model, int used, int limit, int remaining
    ) {}
}

6. 模型响应质量评估与自动切换

java 复制代码
@Service
public class QualityAwareRouter {

    private final ModelRouter router;
    private final Map<String, Double> qualityScores = new ConcurrentHashMap<>();

    public QualityAwareRouter(ModelRouter router) {
        this.router = router;
        qualityScores.put("qwen-turbo", 0.7);
        qualityScores.put("qwen-plus", 0.85);
        qualityScores.put("qwen-coder", 0.9);
        qualityScores.put("deepseek-chat", 0.8);
    }

    /**
     * 反馈驱动的质量评估
     */
    public void recordFeedback(String modelName, double score) {
        // 指数移动平均
        double current = qualityScores.getOrDefault(modelName, 0.5);
        double updated = current * 0.9 + score * 0.1;
        qualityScores.put(modelName, updated);
    }

    /**
     * 结合质量分数和当前路由策略
     */
    public ChatModel routeWithQuality(RouteContext context) {
        ChatModel candidates = router.route(context);
        double score = qualityScores.getOrDefault(getModelName(candidates), 0.5);
        
        // 如果质量分数低于阈值,尝试使用更高的质量模型
        if (score < 0.6) {
            log.info("当前模型质量分数 {} 偏低,切换高分模型", score);
            return getHighestQualityModel();
        }
        
        return candidates;
    }

    private ChatModel getHighestQualityModel() {
        return qualityScores.entrySet().stream()
                .max(Map.Entry.comparingByValue())
                .map(e -> router.getModel(e.getKey()))
                .orElse(null);
    }
}

总结

多模型路由的核心技术点:

  1. ChatModel注册: 多Bean方式配置不同模型
  2. 路由策略: 内容分类→复杂度评估→配额检查→质量评估
  3. 熔断机制: 滑动窗口错误率监控,防止级联故障
  4. 故障转移: 主模型失败自动降级到备用模型
  5. 配额管理: 日维度用量限制,成本可控

参考博客: https://blog.csdn.net/badao_liumang_qizhi

相关推荐
艺杯羹14 分钟前
全栈信创落地实录:基于银河麒麟V10与达梦数据库DM8的SpringBoot工业级适配指南
java·数据库·spring boot·后端·spring
随遇而安zx1 小时前
SpringCloud---Gateway vs Netflix Zuul 网关对比深度解析
spring·spring cloud·gateway
鲨鱼辣钊4 小时前
【FastAPI筑基-Day19】APScheduler定时任务全实战|自动执行、动态启停、后台常驻
java·spring·fastapi
学长毕业设计9 小时前
基于SpringBoot的公益基金管理系统(源码+文档+讲解视频)
java·spring boot·后端
东小西9 小时前
【SAA实战】第 3 篇 · 工具调用全攻略:把业务能力交给 Agent 自己调度
java·后端·spring
东小西9 小时前
【SAA实战】第 4 篇 · Agent 短期记忆:saver 让 Agent 跨轮记得住(threadId 隔离)
java·后端·spring
许彰午9 小时前
22-DataCenter报文序列化
java·低代码·架构·状态模式
2601_962065259 小时前
[MySQL] SQL优化之性能分析
java·sql·mysql
阿kun要赚马内9 小时前
MySQL 索引基础
后端·mysql