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);
}
}
总结
多模型路由的核心技术点:
- ChatModel注册: 多Bean方式配置不同模型
- 路由策略: 内容分类→复杂度评估→配额检查→质量评估
- 熔断机制: 滑动窗口错误率监控,防止级联故障
- 故障转移: 主模型失败自动降级到备用模型
- 配额管理: 日维度用量限制,成本可控