在大模型应用开发中,如何优雅地接入第三方大模型API是每个开发者必须掌握的核心技能。OpenAI的GPT系列和Anthropic的Claude系列是当前最主流的商用大模型,它们的API设计各有特色。本篇文章将详细介绍如何设计一套统一的API接入框架,实现对多种大模型的无缝切换,同时保证系统的稳定性、可扩展性和成本可控。

1.1 主流大模型API概述

OpenAI API以其广泛的模型支持和成熟的开发生态著称。从GPT-3.5到GPT-4o,模型能力持续进化,支持文本生成、函数调用、视觉理解等多模态能力。API采用RESTful风格,通过Bearer Token进行认证,返回格式统一为JSON。

Claude API则以其长上下文窗口和指令遵循能力著称。Claude 3系列支持200K token的超长上下文,在长文档处理、多轮对话等场景有明显优势。Claude API采用不同的消息格式,支持系统提示、用户消息和助手消息的角色区分。

1.2 统一抽象的重要性

在大模型应用开发中,采用统一的抽象层有多重好处。首先是业务逻辑与模型实现解耦,上层业务不需要关心底层调用的是GPT还是Claude。其次是模型切换零感知,当需要更换模型或增加新模型时,不需要修改业务代码。此外,多模型负载均衡成为可能,可以根据模型特点分配不同类型的请求。最后,统一抽象便于实现重试、限流、监控等通用能力。

2.1 核心接口抽象

首先定义统一的模型调用接口:

public interface LlmModelClient {
    /**
     * 同步对话请求
     */
    ChatResponse chat(ChatRequest request);
    /**
     * 流式对话请求
     */
    Flux<String> chatStream(ChatRequest request);
    /**
     * 获取模型元信息
     */
    ModelInfo getModelInfo();
    /**
     * 获取支持的功能特性
     */
    Set<ModelCapability> getCapabilities();
}
/**
 * 对话请求
 */
@Data
public class ChatRequest {
    private String model;
    private String prompt;
    private String systemPrompt;
    private List<Message> history;
    private Double temperature;
    private Integer maxTokens;
    private Double topP;
    private List<String> stop;
    private Map<String, Object> extraParams;
}
/**
 * 对话响应
 */
@Data
public class ChatResponse {
    private String content;
    private String model;
    private String finishReason;
    private Long promptTokens;
    private Long completionTokens;
    private Long totalTokens;
    private Map<String, Object> rawResponse;
}
/**
 * 消息结构
 */
@Data
public class Message {
    private MessageRole role;
    private String content;
    public enum MessageRole {
        SYSTEM, USER, ASSISTANT
    }
}

2.2 OpenAI API实现

OpenAI API的Java客户端实现:

@Component
@RequiredArgsConstructor
class OpenAIClient implements LlmModelClient {
    private final RestTemplate openaiRestTemplate;
    private final ObjectMapper objectMapper;
    private static final String OPENAI_API_URL = "https://api.openai.com/v1/chat/completions";
    private static final String API_KEY = "${OPENAI_API_KEY}";
    @Override
    public ChatResponse chat(ChatRequest request) {
        val openaiRequest = buildOpenAIRequest(request);
        val headers = new HttpHeaders();
        headers.setContentType(MediaType.APPLICATION_JSON);
        headers.setBearerAuth(API_KEY);
        val entity = new HttpEntity<>(openaiRequest, headers);
        try {
            val response = openaiRestTemplate.exchange(
                OPENAI_API_URL,
                HttpMethod.POST,
                entity,
                OpenAIResponse::class.java
            );
            return convertToChatResponse(response.getBody());
        } catch (RestClientException e) {
            throw new LlmApiException("OpenAI API调用失败", e);
        }
    }
    @Override
    public Flux<String> chatStream(ChatRequest request) {
        val openaiRequest = buildOpenAIRequest(request);
        openaiRequest.put("stream", true);
        val headers = new HttpHeaders();
        headers.setContentType(MediaType.APPLICATION_JSON);
        headers.setBearerAuth(API_KEY);
        val entity = new HttpEntity<>(openaiRequest, headers);
        return openaiRestTemplate.exchange(
            OPENAI_API_URL,
            HttpMethod.POST,
            entity,
            String.class
        ).getBody()
        .split("\n")
        .filter(line -> line.startsWith("data: "))
        .filter(line -> !line.equals("data: [DONE]"))
        .map(line -> parseStreamResponse(line));
    }
    private Map<String, Object> buildOpenAIRequest(ChatRequest request) {
        val messages = new ArrayList<Map<String, String>>();
        if (StringUtils.hasText(request.getSystemPrompt())) {
            messages.add(Map.of(
                "role", "system",
                "content", request.getSystemPrompt()
            ));
        }
        if (request.getHistory() != null) {
            request.getHistory().forEach(msg -> {
                val role = switch (msg.getRole()) {
                    case SYSTEM -> "system";
                    case USER -> "user";
                    case ASSISTANT -> "assistant";
                };
                messages.add(Map.of("role", role, "content", msg.getContent()));
            });
        }
        messages.add(Map.of("role", "user", "content", request.getPrompt()));
        val requestBody = new HashMap<String, Object>();
        requestBody.put("model", request.getModel() != null ? request.getModel() : "gpt-3.5-turbo");
        requestBody.put("messages", messages);
        if (request.getTemperature() != null) {
            requestBody.put("temperature", request.getTemperature());
        }
        if (request.getMaxTokens() != null) {
            requestBody.put("max_tokens", request.getMaxTokens());
        }
        if (request.getTopP() != null) {
            requestBody.put("top_p", request.getTopP());
        }
        if (request.getStop() != null && !request.getStop().isEmpty()) {
            requestBody.put("stop", request.getStop());
        }
        return requestBody;
    }
    private ChatResponse convertToChatResponse(OpenAIResponse response) {
        val choice = response.getChoices().get(0);
        val message = choice.getMessage();
        val usage = response.getUsage();
        return ChatResponse.builder()
            .content(message.getContent())
            .model(response.getModel())
            .finishReason(choice.getFinishReason())
            .promptTokens(usage.getPromptTokens())
            .completionTokens(usage.getCompletionTokens())
            .totalTokens(usage.getTotalTokens())
            .build();
    }
    private String parseStreamResponse(String line) {
        val json = line.substring(6); // Remove "data: "
        val chunk = objectMapper.readTree(json);
        return chunk.path("choices")
            .path(0)
            .path("delta")
            .path("content")
            .asText();
    }
    @Override
    public ModelInfo getModelInfo() {
        return ModelInfo.builder()
            .provider("OpenAI")
            .defaultModel("gpt-3.5-turbo")
            .supportedModels(Set.of(
                "gpt-3.5-turbo", "gpt-4", "gpt-4-turbo", "gpt-4o"
            ))
            .maxContextLength(128000)
            .build();
    }
    @Override
    public Set<ModelCapability> getCapabilities() {
        return Set.of(
            ModelCapability.CHAT,
            ModelCapability.STREAMING,
            ModelCapability.FUNCTION_CALLING,
            ModelCapability.VISION
        );
    }
}

2.3 Claude API实现

Claude API的实现需要使用Anthropic特有的请求格式:

@Service
class ClaudeClient implements LlmModelClient {
    private final RestTemplate claudeRestTemplate;
    private final ObjectMapper objectMapper;
    private static final String CLAUDE_API_URL = "https://api.anthropic.com/v1/messages";
    private static final String API_KEY = "${CLAUDE_API_KEY}";
    private static final String CLAUDE_VERSION = "2023-06-01";
    @Override
    public ChatResponse chat(ChatRequest request) {
        val claudeRequest = buildClaudeRequest(request);
        val headers = HttpHeaders();
        headers.setContentType(MediaType.APPLICATION_JSON);
        headers.set("x-api-key", API_KEY);
        headers.set("anthropic-version", CLAUDE_VERSION);
        val entity = new HttpEntity<>(claudeRequest, headers);
        try {
            val response = claudeRestTemplate.exchange(
                CLAUDE_API_URL,
                HttpMethod.POST,
                entity,
                ClaudeResponse::class.java
            );
            return convertToChatResponse(response.getBody());
        } catch (RestClientException e) {
            throw new LlmApiException("Claude API调用失败", e);
        }
    }
    @Override
    public Flux<String> chatStream(ChatRequest request) {
        val claudeRequest = buildClaudeRequest(request);
        claudeRequest.put("stream", true);
        val headers = HttpHeaders();
        headers.setContentType(MediaType.APPLICATION_JSON);
        headers.set("x-api-key", API_KEY);
        headers.set("anthropic-version", CLAUDE_VERSION);
        val entity = new HttpEntity<>(claudeRequest, headers);
        return claudeRestTemplate.exchange(
            CLAUDE_API_URL,
            HttpMethod.POST,
            entity,
            String.class
        ).getBody()
        .split("\n")
        .filter(line -> line.startsWith("data: "))
        .filter(line -> !line.equals("data: [DONE]"))
        .map(line -> parseStreamResponse(line));
    }
    private Map<String, Object> buildClaudeRequest(ChatRequest request) {
        val messages = new ArrayList<Map<String, String>>();
        if (request.getHistory() != null) {
            request.getHistory().forEach(msg -> {
                val role = msg.getRole() == Message.MessageRole.ASSISTANT ? "assistant" : "user";
                messages.add(Map.of("role", role, "content", msg.getContent()));
            });
        }
        messages.add(Map.of("role", "user", "content", request.getPrompt()));
        val requestBody = new HashMap<String, Object>();
        requestBody.put("model", request.getModel() != null ? request.getModel() : "claude-3-sonnet-20240229");
        requestBody.put("messages", messages);
        if (StringUtils.hasText(request.getSystemPrompt())) {
            requestBody.put("system", request.getSystemPrompt());
        }
        if (request.getTemperature() != null) {
            requestBody.put("temperature", request.getTemperature());
        }
        if (request.getMaxTokens() != null) {
            // Claude要求min_tokens >= 1
            requestBody.put("max_tokens", Math.max(request.getMaxTokens(), 1));
        }
        if (request.getTopP() != null) {
            requestBody.put("top_p", request.getTopP());
        }
        if (request.getStop() != null && !request.getStop().isEmpty()) {
            requestBody.put("stop_sequences", request.getStop());
        }
        return requestBody;
    }
    private ChatResponse convertToChatResponse(ClaudeResponse response) {
        val content = response.getContent().get(0);
        val usage = response.getUsage();
        return ChatResponse.builder()
            .content(content.getText())
            .model(response.getModel())
            .finishReason(response.getStopReason())
            .promptTokens(usage.getInputTokens())
            .completionTokens(usage.getOutputTokens())
            .totalTokens(usage.getInputTokens() + usage.getOutputTokens())
            .build();
    }
    @Override
    public ModelInfo getModelInfo() {
        return ModelInfo.builder()
            .provider("Anthropic")
            .defaultModel("claude-3-sonnet-20240229")
            .supportedModels(Set.of(
                "claude-3-opus-20240229",
                "claude-3-sonnet-20240229",
                "claude-3-haiku-20240307"
            ))
            .maxContextLength(200000)
            .build();
    }
    @Override
    public Set<ModelCapability> getCapabilities() {
        return Set.of(
            ModelCapability.CHAT,
            ModelCapability.STREAMING,
            ModelCapability.LONG_CONTEXT
        );
    }
}

3.1 模型路由服务

根据请求特点和模型能力进行智能路由:

@Service
@RequiredArgsConstructor
class ModelRouter {
    private final Map<String, LlmModelClient> clients;
    private final LoadBalancerClient loadBalancer;
    public ChatResponse route(ChatRequest request) {
        val modelId = request.getModel();
        // 1. 能力匹配
        val targetClient = selectClient(modelId, request);
        // 2. 负载均衡(同一模型多实例)
        val client = loadBalance(targetClient);
        // 3. 执行调用
        return executeWithRetry(client, request);
    }
    private LlmModelClient selectClient(String modelId, ChatRequest request) {
        // 优先按指定模型选择
        if (modelId != null) {
            val client = findClientByModel(modelId);
            if (client != null) {
                return client;
            }
        }
        // 根据请求特点自动选择
        if (requiresVision(request)) {
            // 视觉理解优先用GPT-4V
            return findClientByProvider("OpenAI");
        }
        if (requiresLongContext(request)) {
            // 长上下文优先用Claude
            return findClientByProvider("Anthropic");
        }
        // 默认负载均衡
        return loadBalancer.select(clients.values());
    }
    private boolean requiresLongContext(ChatRequest request) {
        // 根据提示词长度或历史消息估算
        val estimatedTokens = estimateTokens(request);
        return estimatedTokens > 32000;
    }
    private LlmModelClient findClientByProvider(String provider) {
        return clients.values().stream()
            .filter(c -> c.getModelInfo().getProvider().equals(provider))
            .findFirst()
            .orElseThrow(() -> new LlmApiException("未找到Provider: " + provider));
    }
}

3.2 智能负载均衡器

@Component
class ModelLoadBalancer {
    private final ConcurrentHashMap<String, AtomicInteger> counters = new ConcurrentHashMap<>();
    private final ConcurrentHashMap<String, Long> lastReset = new ConcurrentHashMap<>();
    public <T> T select(List<T> items, WeightFunction<T> weightFunc) {
        if (items.isEmpty()) {
            throw new IllegalArgumentException("items cannot be empty");
        }
        // 每分钟重置计数器
        resetIfNeeded();
        // 计算加权分数
        val scores = items.stream()
            .mapToObj(item -> new AbstractMap.SimpleEntry<>(
                item,
                weightFunc.apply(item) / (counters.getOrDefault(item.toString(), new AtomicInteger(0)).get() + 1)
            ))
            .toList();
        // 选择分数最高的
        return scores.stream()
            .max(Map.Entry.comparingByValue())
            .map(Map.Entry::getKey)
            .orElse(items.get(0));
    }
    public void recordUsage(String item) {
        counters.computeIfAbsent(item, k -> new AtomicInteger(0)).incrementAndGet();
    }
    private void resetIfNeeded() {
        val now = System.currentTimeMillis();
        lastReset.computeIfAbsent("global", k -> now);
        if (now - lastReset.get("global") > 60000) {
            counters.clear();
            lastReset.put("global", now);
        }
    }
}

4.1 重试机制

@Component
@RequiredArgsConstructor
class ResilientLlmClient implements LlmModelClient {
    private final LlmModelClient delegate;
    private final RetryTemplate retryTemplate;
    @Override
    public ChatResponse chat(ChatRequest request) {
        return retryTemplate.execute(context -> {
            if (context.getRetryCount() > 0) {
                log.warn("LLM API调用重试, attempt: {}, model: {}",
                    context.getRetryCount(), request.getModel());
            }
            return delegate.chat(request);
        });
    }
    @Override
    public Flux<String> chatStream(ChatRequest request) {
        // 流式调用不支持重试,因为已经开始了
        return delegate.chatStream(request)
            .doOnError(error -> log.error("Stream failed", error))
            .onErrorResume(throwable -> Flux.just("[Stream Error: " + throwable.getMessage() + "]"));
    }
}
@Configuration
class RetryConfig {
    @Bean
    public RetryTemplate retryTemplate() {
        val template = new RetryTemplate();
        val policy = new CompositeRetryPolicy();
        policy.setPolicies(new RetryPolicy[]{
            new MapRetryPolicy(Map.of(
                IOException.class, true,
                LlmApiException.class, e -> ((LlmApiException) e).isRetryable()
            )),
            new MaxAttemptsRetryPolicy(3),
            new TimeoutRetryPolicy(30000)
        });
        template.setRetryPolicy(policy);
        template.setBackOffPolicy(new ExponentialBackOffPolicy(
            1000L, 2.0, 10000L
        ));
        return template;
    }
}

4.2 熔断降级

@Service
class CircuitBreakerLlmClient implements LlmModelClient {
    private final LlmModelClient delegate;
    private final AtomicReference<State> state = new AtomicReference<>(State.CLOSED);
    private final AtomicInteger failureCount = new AtomicInteger(0);
    private final AtomicInteger successCount = new AtomicInteger(0);
    private volatile long lastFailureTime = 0;
    private static final int FAILURE_THRESHOLD = 5;
    private static final int SUCCESS_THRESHOLD = 3;
    private static final long RECOVERY_TIMEOUT_MS = 60000;
    public enum State { CLOSED, OPEN, HALF_OPEN }
    @Override
    public ChatResponse chat(ChatRequest request) {
        if (state.get() == State.OPEN) {
            if (System.currentTimeMillis() - lastFailureTime > RECOVERY_TIMEOUT_MS) {
                state.compareAndSet(State.OPEN, State.HALF_OPEN);
            } else {
                throw new LlmApiException("Circuit breaker OPEN, fallback triggered");
            }
        }
        try {
            val response = delegate.chat(request);
            onSuccess();
            return response;
        } catch (Exception e) {
            onFailure(e);
            throw e;
        }
    }
    private void onSuccess() {
        failureCount.set(0);
        if (state.compareAndSet(State.HALF_OPEN, State.CLOSED)) {
            log.info("Circuit breaker CLOSED after recovery");
        }
    }
    private void onFailure(Exception e) {
        lastFailureTime = System.currentTimeMillis();
        val failures = failureCount.incrementAndGet();
        if (failures >= FAILURE_THRESHOLD) {
            if (state.compareAndSet(State.CLOSED, State.OPEN)) {
                log.error("Circuit breaker OPEN due to {} consecutive failures", failures);
            }
        }
    }
}

4.3 限流与配额管理

@Component
public class RateLimiter {
    private final ConcurrentHashMap<String, RateLimiter> userLimiters = new ConcurrentHashMap<>();
    private final ConcurrentHashMap<String, TokenBucket> buckets = new ConcurrentHashMap<>();
    private static final int DEFAULT_TOKENS = 60;
    private static final long REFILL_INTERVAL_MS = 60000;
    public boolean tryAcquire(String userId, String model) {
        val key = userId + ":" + model;
        val bucket = buckets.computeIfAbsent(key, k -> new TokenBucket(DEFAULT_TOKENS));
        synchronized (bucket) {
            if (bucket.tryConsume()) {
                return true;
            }
        }
        log.warn("Rate limit exceeded for user: {}, model: {}", userId, model);
        return false;
    }
    public static class TokenBucket {
        private final int capacity;
        private volatile int tokens;
        private volatile long lastRefillTime;
        public TokenBucket(int capacity) {
            this.capacity = capacity;
            this.tokens = capacity;
            this.lastRefillTime = System.currentTimeMillis();
        }
        public synchronized boolean tryConsume() {
            refill();
            if (tokens > 0) {
                tokens--;
                return true;
            }
            return false;
        }
        private void refill() {
            val now = System.currentTimeMillis();
            val elapsed = now - lastRefillTime;
            if (elapsed >= REFILL_INTERVAL_MS) {
                tokens = capacity;
                lastRefillTime = now;
            }
        }
    }
}

5.1 用量追踪

@Service
@RequiredArgsConstructor
public class UsageTracker {
    private final MetricsFacade metrics;
    private final UsageRepository usageRepository;
    public void record(ChatRequest request, ChatResponse response, String userId) {
        val usage = LlmUsage.builder()
            .userId(userId)
            .model(response.getModel())
            .provider(determineProvider(response.getModel()))
            .promptTokens(response.getPromptTokens())
            .completionTokens(response.getCompletionTokens())
            .totalTokens(response.getTotalTokens())
            .cost(calculateCost(response))
            .timestamp(System.currentTimeMillis())
            .build();
        usageRepository.save(usage);
        // 实时指标
        metrics.increment("llm.requests.total", Tags.of("model", response.getModel()));
        metrics.gauge("llm.tokens.total", response.getTotalTokens(),
            Tags.of("model", response.getModel(), "type", "completion"));
    }
    private BigDecimal calculateCost(ChatResponse response) {
        val model = response.getModel();
        val promptTokens = response.getPromptTokens();
        val completionTokens = response.getCompletionTokens();
        // OpenAI定价(示例)
        if (model.startsWith("gpt-4")) {
            return BigDecimal.valueOf(0.03)
                .multiply(BigDecimal.valueOf(promptTokens))
                .add(BigDecimal.valueOf(0.06)
                    .multiply(BigDecimal.valueOf(completionTokens)))
                .divide(BigDecimal.valueOf(1000), 4, RoundingMode.HALF_UP);
        } else if (model.startsWith("gpt-3.5")) {
            return BigDecimal.valueOf(0.0015)
                .multiply(BigDecimal.valueOf(promptTokens))
                .add(BigDecimal.valueOf(0.002)
                    .multiply(BigDecimal.valueOf(completionTokens)))
                .divide(BigDecimal.valueOf(1000), 4, RoundingMode.HALF_UP);
        }
        // Claude定价(示例)
        if (model.contains("claude-3-opus")) {
            return BigDecimal.valueOf(0.015)
                .multiply(BigDecimal.valueOf(promptTokens))
                .add(BigDecimal.valueOf(0.075)
                    .multiply(BigDecimal.valueOf(completionTokens)))
                .divide(BigDecimal.valueOf(1000), 4, RoundingMode.HALF_UP);
        }
        return BigDecimal.ZERO;
    }
}

5.2 成本仪表盘

@RestController
@RequestMapping("/api/admin/costs")
@RequiredArgsConstructor
public class CostDashboardController {
    private final UsageRepository usageRepository;
    @GetMapping("/summary")
    public CostSummary getSummary(
        @RequestParam(required = false) String startDate,
        @RequestParam(required = false) String endDate
    ) {
        val start = parseDate(startDate, DateUtils.addDays(new Date(), -30));
        val end = parseDate(endDate, new Date());
        val usageList = usageRepository.findByDateRange(start, end);
        return CostSummary.builder()
            .totalRequests(usageList.size())
            .totalPromptTokens(usageList.stream().mapToLong(LlmUsage::getPromptTokens).sum())
            .totalCompletionTokens(usageList.stream().mapToLong(LlmUsage::getCompletionTokens).sum())
            .totalCost(usageList.stream().map(LlmUsage::getCost).reduce(BigDecimal.ZERO, BigDecimal::add))
            .costByModel(usageList.stream()
                .collect(groupingBy(LlmUsage::getModel,
                    reducing(BigDecimal.ZERO, LlmUsage::getCost, BigDecimal::add))))
            .costByDay(usageList.stream()
                .collect(groupingBy(
                    u -> DateFormatUtils.format(u.getTimestamp(), "yyyy-MM-dd"),
                    reducing(BigDecimal.ZERO, LlmUsage::getCost, BigDecimal::add))))
            .build();
    }
}

本章详细介绍了OpenAI和Claude API的接入设计,从接口抽象到具体实现,从路由策略到稳定性保障,完整覆盖了大模型API接入的核心知识点。

通过统一的接口抽象和智能路由,系统可以同时支持多种大模型,根据业务场景灵活选择最适合的模型。通过重试、熔断、限流等机制,确保系统在高并发和异常情况下的稳定性。通过用量追踪和成本监控,实现了大模型调用的精细化管理。

下一章将介绍LangChain4j微服务集成,帮助开发者在大模型应用中快速构建复杂的业务逻辑。

Logo

有“AI”的1024 = 2048,欢迎大家加入2048 AI社区

更多推荐