feat(providers): add Cloudflare AI Gateway support
Implements Cloudflare AI Gateway as an OpenAI-compatible provider with support for: - Unified Billing mode (cf_token only) - BYOK (Bring Your Own Key) mode - Cloudflare Workers AI serverless models Includes provider implementation, configuration support, comprehensive tests, documentation updates, and config examples.
This commit is contained in:
parent
2c8416e658
commit
8d09a2b3e2
11 changed files with 1103 additions and 659 deletions
43
README.zh.md
43
README.zh.md
|
|
@ -464,6 +464,7 @@ Agent 读取 HEARTBEAT.md
|
|||
| **神算云** | `shengsuanyun/` | `https://router.shengsuanyun.com/api/v1` | OpenAI | - |
|
||||
| **Antigravity** | `antigravity/` | Google Cloud | 自定义 | 仅 OAuth |
|
||||
| **GitHub Copilot** | `github-copilot/` | `localhost:4321` | gRPC | - |
|
||||
| **Cloudflare AI Gateway** | `cloudflare/` | `https://gateway.ai.cloudflare.com/v1/YOUR_ACCOUNT_ID/YOUR_GATEWAY_ID/compat` | OpenAI | [获取令牌](https://www.cloudflare.com/zh-tw/developer-platform/products/ai-gateway/) |
|
||||
|
||||
#### 基础配置示例
|
||||
|
||||
|
|
@ -559,6 +560,48 @@ Agent 读取 HEARTBEAT.md
|
|||
}
|
||||
```
|
||||
|
||||
**Cloudflare AI Gateway(统一计费)**
|
||||
|
||||
使用 Cloudflare 的计费来支付上游服务商 —— 无需各个服务商的 API Key:
|
||||
|
||||
```json
|
||||
{
|
||||
"model_name": "cf-gpt5",
|
||||
"model": "cloudflare/openai/gpt-5.2",
|
||||
"cf_token": "YOUR_CLOUDFLARE_AIG_TOKEN",
|
||||
"api_base": "https://gateway.ai.cloudflare.com/v1/YOUR_ACCOUNT_ID/YOUR_GATEWAY_ID/compat"
|
||||
}
|
||||
```
|
||||
|
||||
**Cloudflare AI Gateway(BYOK — 自带密钥)**
|
||||
|
||||
通过 Cloudflare 路由请求,同时使用自己的服务商 API Key:
|
||||
|
||||
```json
|
||||
{
|
||||
"model_name": "cf-claude",
|
||||
"model": "cloudflare/anthropic/claude-sonnet-4.6",
|
||||
"api_key": "sk-ant-your-key",
|
||||
"cf_token": "YOUR_CLOUDFLARE_AIG_TOKEN",
|
||||
"api_base": "https://gateway.ai.cloudflare.com/v1/YOUR_ACCOUNT_ID/YOUR_GATEWAY_ID/compat"
|
||||
}
|
||||
```
|
||||
|
||||
**Cloudflare Workers AI**
|
||||
|
||||
通过 Cloudflare Workers AI 运行无服务器推理模型:
|
||||
|
||||
```json
|
||||
{
|
||||
"model_name": "cf-gpt-oss-120b",
|
||||
"model": "cloudflare/workers-ai/@cf/openai/gpt-oss-120b",
|
||||
"cf_token": "YOUR_CLOUDFLARE_AIG_TOKEN",
|
||||
"api_base": "https://gateway.ai.cloudflare.com/v1/YOUR_ACCOUNT_ID/YOUR_GATEWAY_ID/compat"
|
||||
}
|
||||
```
|
||||
|
||||
> 详见 [Cloudflare AI Gateway 文档](https://developers.cloudflare.com/ai-gateway/) 和 [Workers AI 模型列表](https://developers.cloudflare.com/workers-ai/models/)。
|
||||
|
||||
#### 负载均衡
|
||||
|
||||
为同一个模型名称配置多个端点——PicoClaw 会自动在它们之间轮询:
|
||||
|
|
|
|||
|
|
@ -43,6 +43,12 @@
|
|||
"model": "openai/gpt-5.2",
|
||||
"api_key": "sk-key2",
|
||||
"api_base": "https://api2.example.com/v1"
|
||||
},
|
||||
{
|
||||
"model_name": "cf-gpt-oss-120b",
|
||||
"model": "cloudflare/workers-ai/@cf/openai/gpt-oss-120b",
|
||||
"cf_token": "CLOUDFLARE_AIG_TOKEN",
|
||||
"api_base": "https://gateway.ai.cloudflare.com/v1/YOUR_ACCOUNT_ID/YOUR_GATEWAY_ID/compat"
|
||||
}
|
||||
],
|
||||
"channels": {
|
||||
|
|
|
|||
|
|
@ -405,6 +405,9 @@ type ModelConfig struct {
|
|||
ConnectMode string `json:"connect_mode,omitempty"` // Connection mode: stdio, grpc
|
||||
Workspace string `json:"workspace,omitempty"` // Workspace path for CLI-based providers
|
||||
|
||||
// Cloudflare AI Gateway
|
||||
CfToken string `json:"cf_token,omitempty"` // Cloudflare AI Gateway token for cf-aig-authorization header
|
||||
|
||||
// Optional optimizations
|
||||
RPM int `json:"rpm,omitempty"` // Requests per minute limit
|
||||
MaxTokensField string `json:"max_tokens_field,omitempty"` // Field name for max tokens (e.g., "max_completion_tokens")
|
||||
|
|
|
|||
|
|
@ -270,6 +270,18 @@ func DefaultConfig() *Config {
|
|||
APIBase: "http://localhost:8000/v1",
|
||||
APIKey: "",
|
||||
},
|
||||
|
||||
// Cloudflare AI Gateway - https://developers.cloudflare.com/ai-gateway/
|
||||
// Supports Unified Billing (cf_token only) or BYOK (add api_key).
|
||||
// Model format: cloudflare/{provider}/{model}
|
||||
// e.g., cloudflare/openai/gpt-5.2, cloudflare/anthropic/claude-sonnet-4.6
|
||||
// cloudflare/workers-ai/@cf/openai/gpt-oss-120b
|
||||
{
|
||||
ModelName: "cf-gpt-oss-120b",
|
||||
Model: "cloudflare/workers-ai/@cf/openai/gpt-oss-120b",
|
||||
APIBase: "https://gateway.ai.cloudflare.com/v1/YOUR_ACCOUNT_ID/YOUR_GATEWAY_ID/compat",
|
||||
CfToken: "",
|
||||
},
|
||||
},
|
||||
Gateway: GatewayConfig{
|
||||
Host: "127.0.0.1",
|
||||
|
|
|
|||
99
pkg/providers/cloudflare_provider.go
Normal file
99
pkg/providers/cloudflare_provider.go
Normal file
|
|
@ -0,0 +1,99 @@
|
|||
// CloudflareProvider is an LLM provider for Cloudflare AI Gateway.
|
||||
//
|
||||
// Cloudflare AI Gateway provides a unified OpenAI-compatible endpoint that
|
||||
// proxies requests to multiple upstream AI providers (OpenAI, Anthropic,
|
||||
// Google, Workers AI, etc.) through a single API.
|
||||
//
|
||||
// Endpoint format:
|
||||
//
|
||||
// https://gateway.ai.cloudflare.com/v1/{account_id}/{gateway_id}/compat/chat/completions
|
||||
//
|
||||
// Authentication modes:
|
||||
// - Unified Billing: only cf_token required (Cloudflare pays the upstream provider)
|
||||
// - BYOK (Bring Your Own Key): api_key for the upstream provider + optional cf_token
|
||||
// for authenticated gateways
|
||||
//
|
||||
// Model format in the request body: "{provider}/{model}"
|
||||
//
|
||||
// e.g., "openai/gpt-5.2", "anthropic/claude-sonnet-4.6",
|
||||
// "workers-ai/@cf/openai/gpt-oss-120b"
|
||||
//
|
||||
// References:
|
||||
// - https://developers.cloudflare.com/ai-gateway/usage/chat-completion/
|
||||
// - https://developers.cloudflare.com/ai-gateway/features/unified-billing/
|
||||
// - https://developers.cloudflare.com/ai-gateway/configuration/authentication/
|
||||
// - https://developers.cloudflare.com/workers-ai/models/
|
||||
|
||||
package providers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/providers/openai_compat"
|
||||
)
|
||||
|
||||
type CloudflareProvider struct {
|
||||
delegate *openai_compat.Provider
|
||||
}
|
||||
|
||||
// CfAIGAuthHeader is the HTTP header name used by Cloudflare AI Gateway
|
||||
// for gateway-level authentication.
|
||||
const CfAIGAuthHeader = "cf-aig-authorization"
|
||||
|
||||
// NewCloudflareProvider creates a new Cloudflare AI Gateway provider.
|
||||
//
|
||||
// Parameters:
|
||||
// - apiKey: upstream provider API key (for BYOK mode), sent as Authorization header.
|
||||
// Leave empty for Unified Billing mode.
|
||||
// - apiBase: the AI Gateway compat endpoint URL, e.g.,
|
||||
// "https://gateway.ai.cloudflare.com/v1/{account_id}/{gateway_id}/compat"
|
||||
// - cfToken: Cloudflare API token for gateway authentication, sent as
|
||||
// cf-aig-authorization header. Required for Unified Billing; optional for BYOK
|
||||
// (only needed if the gateway has authentication enabled).
|
||||
// - proxy: optional HTTP proxy URL.
|
||||
// - maxTokensField: optional field name override for max tokens.
|
||||
// - requestTimeoutSeconds: request timeout in seconds (0 for default).
|
||||
func NewCloudflareProvider(
|
||||
apiKey, apiBase, cfToken, proxy, maxTokensField string,
|
||||
requestTimeoutSeconds int,
|
||||
) (*CloudflareProvider, error) {
|
||||
if apiBase == "" {
|
||||
return nil, fmt.Errorf("api_base is required for Cloudflare AI Gateway")
|
||||
}
|
||||
|
||||
// Build extra headers for Cloudflare AI Gateway authentication
|
||||
extraHeaders := make(map[string]string)
|
||||
if cfToken != "" {
|
||||
extraHeaders[CfAIGAuthHeader] = "Bearer " + cfToken
|
||||
}
|
||||
|
||||
opts := []openai_compat.Option{
|
||||
openai_compat.WithMaxTokensField(maxTokensField),
|
||||
openai_compat.WithExtraHeaders(extraHeaders),
|
||||
}
|
||||
if requestTimeoutSeconds > 0 {
|
||||
opts = append(opts, openai_compat.WithRequestTimeout(
|
||||
time.Duration(requestTimeoutSeconds)*time.Second,
|
||||
))
|
||||
}
|
||||
|
||||
delegate := openai_compat.NewProvider(apiKey, apiBase, proxy, opts...)
|
||||
|
||||
return &CloudflareProvider{delegate: delegate}, nil
|
||||
}
|
||||
|
||||
func (p *CloudflareProvider) Chat(
|
||||
ctx context.Context,
|
||||
messages []Message,
|
||||
tools []ToolDefinition,
|
||||
model string,
|
||||
options map[string]any,
|
||||
) (*LLMResponse, error) {
|
||||
return p.delegate.Chat(ctx, messages, tools, model, options)
|
||||
}
|
||||
|
||||
func (p *CloudflareProvider) GetDefaultModel() string {
|
||||
return ""
|
||||
}
|
||||
394
pkg/providers/cloudflare_provider_test.go
Normal file
394
pkg/providers/cloudflare_provider_test.go
Normal file
|
|
@ -0,0 +1,394 @@
|
|||
// PicoClaw - Ultra-lightweight personal AI agent
|
||||
// License: MIT
|
||||
//
|
||||
// Copyright (c) 2026 PicoClaw contributors
|
||||
|
||||
package providers
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestNewCloudflareProvider_RequiresAPIBase(t *testing.T) {
|
||||
_, err := NewCloudflareProvider("key", "", "cf-token", "", "", 0)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for empty api_base")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "api_base is required") {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewCloudflareProvider_UnifiedBilling(t *testing.T) {
|
||||
// Unified Billing mode: cf_token only, no api_key
|
||||
provider, err := NewCloudflareProvider(
|
||||
"",
|
||||
"https://gateway.ai.cloudflare.com/v1/acct/gw/compat",
|
||||
"cf-test-token",
|
||||
"", "", 0,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("NewCloudflareProvider() error = %v", err)
|
||||
}
|
||||
if provider == nil {
|
||||
t.Fatal("expected non-nil provider")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewCloudflareProvider_BYOK(t *testing.T) {
|
||||
// BYOK mode: api_key + cf_token
|
||||
provider, err := NewCloudflareProvider(
|
||||
"sk-upstream-key",
|
||||
"https://gateway.ai.cloudflare.com/v1/acct/gw/compat",
|
||||
"cf-test-token",
|
||||
"", "", 0,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("NewCloudflareProvider() error = %v", err)
|
||||
}
|
||||
if provider == nil {
|
||||
t.Fatal("expected non-nil provider")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCloudflareProvider_Chat_SendsCfAigAuthHeader(t *testing.T) {
|
||||
var receivedHeaders http.Header
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
receivedHeaders = r.Header.Clone()
|
||||
resp := map[string]any{
|
||||
"choices": []map[string]any{
|
||||
{
|
||||
"message": map[string]any{"content": "Hello from Cloudflare!"},
|
||||
"finish_reason": "stop",
|
||||
},
|
||||
},
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(resp)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
provider, err := NewCloudflareProvider("", server.URL, "cf-my-token", "", "", 0)
|
||||
if err != nil {
|
||||
t.Fatalf("NewCloudflareProvider() error = %v", err)
|
||||
}
|
||||
|
||||
result, err := provider.Chat(
|
||||
t.Context(),
|
||||
[]Message{{Role: "user", Content: "hi"}},
|
||||
nil,
|
||||
"openai/gpt-5.2",
|
||||
nil,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("Chat() error = %v", err)
|
||||
}
|
||||
|
||||
// Verify cf-aig-authorization header was sent
|
||||
cfAuth := receivedHeaders.Get(CfAIGAuthHeader)
|
||||
if cfAuth != "Bearer cf-my-token" {
|
||||
t.Errorf("cf-aig-authorization header = %q, want %q", cfAuth, "Bearer cf-my-token")
|
||||
}
|
||||
|
||||
// Verify no Authorization header is set when api_key is empty
|
||||
authHeader := receivedHeaders.Get("Authorization")
|
||||
if authHeader != "" {
|
||||
t.Errorf("Authorization header should be empty for Unified Billing, got %q", authHeader)
|
||||
}
|
||||
|
||||
if result.Content != "Hello from Cloudflare!" {
|
||||
t.Errorf("Content = %q, want %q", result.Content, "Hello from Cloudflare!")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCloudflareProvider_Chat_BYOK_SendsBothHeaders(t *testing.T) {
|
||||
var receivedHeaders http.Header
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
receivedHeaders = r.Header.Clone()
|
||||
resp := map[string]any{
|
||||
"choices": []map[string]any{
|
||||
{
|
||||
"message": map[string]any{"content": "ok"},
|
||||
"finish_reason": "stop",
|
||||
},
|
||||
},
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(resp)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
provider, err := NewCloudflareProvider(
|
||||
"sk-upstream-key",
|
||||
server.URL,
|
||||
"cf-my-token",
|
||||
"", "", 0,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("NewCloudflareProvider() error = %v", err)
|
||||
}
|
||||
|
||||
_, err = provider.Chat(
|
||||
t.Context(),
|
||||
[]Message{{Role: "user", Content: "hi"}},
|
||||
nil,
|
||||
"anthropic/claude-sonnet-4.6",
|
||||
nil,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("Chat() error = %v", err)
|
||||
}
|
||||
|
||||
// Verify both headers are sent for BYOK mode
|
||||
cfAuth := receivedHeaders.Get(CfAIGAuthHeader)
|
||||
if cfAuth != "Bearer cf-my-token" {
|
||||
t.Errorf("cf-aig-authorization = %q, want %q", cfAuth, "Bearer cf-my-token")
|
||||
}
|
||||
|
||||
authHeader := receivedHeaders.Get("Authorization")
|
||||
if authHeader != "Bearer sk-upstream-key" {
|
||||
t.Errorf("Authorization = %q, want %q", authHeader, "Bearer sk-upstream-key")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCloudflareProvider_Chat_ModelPassthrough(t *testing.T) {
|
||||
// Verify the model string (e.g., "openai/gpt-5.2") is sent as-is in the request body
|
||||
var requestBody map[string]any
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if err := json.NewDecoder(r.Body).Decode(&requestBody); err != nil {
|
||||
http.Error(w, err.Error(), http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
resp := map[string]any{
|
||||
"choices": []map[string]any{
|
||||
{
|
||||
"message": map[string]any{"content": "ok"},
|
||||
"finish_reason": "stop",
|
||||
},
|
||||
},
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(resp)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
provider, err := NewCloudflareProvider("", server.URL, "cf-token", "", "", 0)
|
||||
if err != nil {
|
||||
t.Fatalf("NewCloudflareProvider() error = %v", err)
|
||||
}
|
||||
|
||||
_, err = provider.Chat(
|
||||
t.Context(),
|
||||
[]Message{{Role: "user", Content: "hi"}},
|
||||
nil,
|
||||
"openai/gpt-5.2",
|
||||
nil,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("Chat() error = %v", err)
|
||||
}
|
||||
|
||||
// Cloudflare expects the full provider/model string
|
||||
if requestBody["model"] != "openai/gpt-5.2" {
|
||||
t.Errorf("model = %v, want %q", requestBody["model"], "openai/gpt-5.2")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCloudflareProvider_Chat_WorkersAI(t *testing.T) {
|
||||
var requestBody map[string]any
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if err := json.NewDecoder(r.Body).Decode(&requestBody); err != nil {
|
||||
http.Error(w, err.Error(), http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
resp := map[string]any{
|
||||
"choices": []map[string]any{
|
||||
{
|
||||
"message": map[string]any{"content": "Workers AI response"},
|
||||
"finish_reason": "stop",
|
||||
},
|
||||
},
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(resp)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
provider, err := NewCloudflareProvider("", server.URL, "cf-token", "", "", 0)
|
||||
if err != nil {
|
||||
t.Fatalf("NewCloudflareProvider() error = %v", err)
|
||||
}
|
||||
|
||||
_, err = provider.Chat(
|
||||
t.Context(),
|
||||
[]Message{{Role: "user", Content: "hi"}},
|
||||
nil,
|
||||
"workers-ai/@cf/meta/llama-3.3-70b-instruct-fp8-fast",
|
||||
nil,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("Chat() error = %v", err)
|
||||
}
|
||||
|
||||
// Workers AI model should be passed through as-is
|
||||
if requestBody["model"] != "workers-ai/@cf/meta/llama-3.3-70b-instruct-fp8-fast" {
|
||||
t.Errorf("model = %v, want %q", requestBody["model"], "workers-ai/@cf/meta/llama-3.3-70b-instruct-fp8-fast")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCloudflareProvider_Chat_NoCfToken(t *testing.T) {
|
||||
// When cf_token is empty, no cf-aig-authorization header should be sent
|
||||
var receivedHeaders http.Header
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
receivedHeaders = r.Header.Clone()
|
||||
resp := map[string]any{
|
||||
"choices": []map[string]any{
|
||||
{
|
||||
"message": map[string]any{"content": "ok"},
|
||||
"finish_reason": "stop",
|
||||
},
|
||||
},
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(resp)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
provider, err := NewCloudflareProvider("sk-my-key", server.URL, "", "", "", 0)
|
||||
if err != nil {
|
||||
t.Fatalf("NewCloudflareProvider() error = %v", err)
|
||||
}
|
||||
|
||||
_, err = provider.Chat(
|
||||
t.Context(),
|
||||
[]Message{{Role: "user", Content: "hi"}},
|
||||
nil,
|
||||
"openai/gpt-5.2",
|
||||
nil,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("Chat() error = %v", err)
|
||||
}
|
||||
|
||||
// No cf-aig-authorization header when cf_token is empty
|
||||
cfAuth := receivedHeaders.Get(CfAIGAuthHeader)
|
||||
if cfAuth != "" {
|
||||
t.Errorf("cf-aig-authorization should be empty, got %q", cfAuth)
|
||||
}
|
||||
|
||||
// Authorization header should still be set with the api_key
|
||||
authHeader := receivedHeaders.Get("Authorization")
|
||||
if authHeader != "Bearer sk-my-key" {
|
||||
t.Errorf("Authorization = %q, want %q", authHeader, "Bearer sk-my-key")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCloudflareProvider_Chat_ToolCalls(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
resp := map[string]any{
|
||||
"choices": []map[string]any{
|
||||
{
|
||||
"message": map[string]any{
|
||||
"content": "",
|
||||
"tool_calls": []map[string]any{
|
||||
{
|
||||
"id": "call_cf1",
|
||||
"type": "function",
|
||||
"function": map[string]any{
|
||||
"name": "get_weather",
|
||||
"arguments": `{"city":"London"}`,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
"finish_reason": "tool_calls",
|
||||
},
|
||||
},
|
||||
"usage": map[string]any{
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 5,
|
||||
"total_tokens": 15,
|
||||
},
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(resp)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
provider, err := NewCloudflareProvider("", server.URL, "cf-token", "", "", 0)
|
||||
if err != nil {
|
||||
t.Fatalf("NewCloudflareProvider() error = %v", err)
|
||||
}
|
||||
|
||||
result, err := provider.Chat(
|
||||
t.Context(),
|
||||
[]Message{{Role: "user", Content: "What's the weather?"}},
|
||||
nil,
|
||||
"openai/gpt-5.2",
|
||||
nil,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("Chat() error = %v", err)
|
||||
}
|
||||
|
||||
if len(result.ToolCalls) != 1 {
|
||||
t.Fatalf("len(ToolCalls) = %d, want 1", len(result.ToolCalls))
|
||||
}
|
||||
if result.ToolCalls[0].Name != "get_weather" {
|
||||
t.Errorf("ToolCalls[0].Name = %q, want %q", result.ToolCalls[0].Name, "get_weather")
|
||||
}
|
||||
if result.ToolCalls[0].Arguments["city"] != "London" {
|
||||
t.Errorf("ToolCalls[0].Arguments[city] = %v, want London", result.ToolCalls[0].Arguments["city"])
|
||||
}
|
||||
if result.Usage == nil {
|
||||
t.Fatal("expected non-nil usage")
|
||||
}
|
||||
if result.Usage.TotalTokens != 15 {
|
||||
t.Errorf("Usage.TotalTokens = %d, want 15", result.Usage.TotalTokens)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCloudflareProvider_Chat_HTTPError(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
http.Error(w, `{"error":"rate limited"}`, http.StatusTooManyRequests)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
provider, err := NewCloudflareProvider("", server.URL, "cf-token", "", "", 0)
|
||||
if err != nil {
|
||||
t.Fatalf("NewCloudflareProvider() error = %v", err)
|
||||
}
|
||||
|
||||
_, err = provider.Chat(
|
||||
t.Context(),
|
||||
[]Message{{Role: "user", Content: "hi"}},
|
||||
nil,
|
||||
"openai/gpt-5.2",
|
||||
nil,
|
||||
)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for HTTP 429")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "429") {
|
||||
t.Errorf("error should contain status code 429, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCloudflareProvider_GetDefaultModel(t *testing.T) {
|
||||
provider, err := NewCloudflareProvider("", "https://example.com", "token", "", "", 0)
|
||||
if err != nil {
|
||||
t.Fatalf("NewCloudflareProvider() error = %v", err)
|
||||
}
|
||||
if model := provider.GetDefaultModel(); model != "" {
|
||||
t.Errorf("GetDefaultModel() = %q, want empty string", model)
|
||||
}
|
||||
}
|
||||
|
|
@ -53,7 +53,7 @@ func ExtractProtocol(model string) (protocol, modelID string) {
|
|||
|
||||
// CreateProviderFromConfig creates a provider based on the ModelConfig.
|
||||
// It uses the protocol prefix in the Model field to determine which provider to create.
|
||||
// Supported protocols: openai, anthropic, antigravity, claude-cli, codex-cli, github-copilot
|
||||
// Supported protocols: openai, anthropic, antigravity, claude-cli, codex-cli, github-copilot, cloudflare
|
||||
// Returns the provider, the model ID (without protocol prefix), and any error.
|
||||
func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, error) {
|
||||
if cfg == nil {
|
||||
|
|
@ -139,6 +139,36 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
|
|||
case "antigravity":
|
||||
return NewAntigravityProvider(), modelID, nil
|
||||
|
||||
case "cloudflare":
|
||||
// Cloudflare AI Gateway (OpenAI-compatible unified endpoint).
|
||||
// Model format: cloudflare/{upstream_provider}/{model}
|
||||
// e.g., "cloudflare/openai/gpt-5.2" → modelID = "openai/gpt-5.2"
|
||||
// e.g., "cloudflare/workers-ai/@cf/openai/gpt-oss-120b" → modelID = "workers-ai/@cf/openai/gpt-oss-120b"
|
||||
//
|
||||
// Auth modes:
|
||||
// - Unified Billing: cf_token only (no api_key needed)
|
||||
// - BYOK: api_key (upstream provider key) + optional cf_token
|
||||
if cfg.APIBase == "" {
|
||||
return nil, "", fmt.Errorf(
|
||||
"api_base is required for cloudflare protocol (e.g., https://gateway.ai.cloudflare.com/v1/ACCOUNT_ID/GATEWAY_ID/compat)",
|
||||
)
|
||||
}
|
||||
if cfg.CfToken == "" && cfg.APIKey == "" {
|
||||
return nil, "", fmt.Errorf("cf_token or api_key is required for cloudflare protocol")
|
||||
}
|
||||
provider, err := NewCloudflareProvider(
|
||||
cfg.APIKey,
|
||||
cfg.APIBase,
|
||||
cfg.CfToken,
|
||||
cfg.Proxy,
|
||||
cfg.MaxTokensField,
|
||||
cfg.RequestTimeout,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf("creating cloudflare provider: %w", err)
|
||||
}
|
||||
return provider, modelID, nil
|
||||
|
||||
case "claude-cli", "claudecli":
|
||||
workspace := cfg.Workspace
|
||||
if workspace == "" {
|
||||
|
|
|
|||
|
|
@ -64,6 +64,24 @@ func TestExtractProtocol(t *testing.T) {
|
|||
wantProtocol: "nvidia",
|
||||
wantModelID: "meta/llama-3.1-8b",
|
||||
},
|
||||
{
|
||||
name: "cloudflare openai upstream",
|
||||
model: "cloudflare/openai/gpt-5.2",
|
||||
wantProtocol: "cloudflare",
|
||||
wantModelID: "openai/gpt-5.2",
|
||||
},
|
||||
{
|
||||
name: "cloudflare anthropic upstream",
|
||||
model: "cloudflare/anthropic/claude-sonnet-4.6",
|
||||
wantProtocol: "cloudflare",
|
||||
wantModelID: "anthropic/claude-sonnet-4.6",
|
||||
},
|
||||
{
|
||||
name: "cloudflare workers-ai deep path",
|
||||
model: "cloudflare/workers-ai/@cf/meta/llama-3.3-70b-instruct-fp8-fast",
|
||||
wantProtocol: "cloudflare",
|
||||
wantModelID: "workers-ai/@cf/meta/llama-3.3-70b-instruct-fp8-fast",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
|
|
@ -220,6 +238,106 @@ func TestCreateProviderFromConfig_MissingAPIKey(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestCreateProviderFromConfig_Cloudflare_UnifiedBilling(t *testing.T) {
|
||||
cfg := &config.ModelConfig{
|
||||
ModelName: "cf-gpt5",
|
||||
Model: "cloudflare/openai/gpt-5.2",
|
||||
CfToken: "cf-test-token",
|
||||
APIBase: "https://gateway.ai.cloudflare.com/v1/acct/gw/compat",
|
||||
}
|
||||
|
||||
provider, modelID, err := CreateProviderFromConfig(cfg)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateProviderFromConfig() error = %v", err)
|
||||
}
|
||||
if provider == nil {
|
||||
t.Fatal("expected non-nil provider")
|
||||
}
|
||||
if _, ok := provider.(*CloudflareProvider); !ok {
|
||||
t.Fatalf("expected *CloudflareProvider, got %T", provider)
|
||||
}
|
||||
// modelID should be everything after "cloudflare/"
|
||||
if modelID != "openai/gpt-5.2" {
|
||||
t.Errorf("modelID = %q, want %q", modelID, "openai/gpt-5.2")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateProviderFromConfig_Cloudflare_BYOK(t *testing.T) {
|
||||
cfg := &config.ModelConfig{
|
||||
ModelName: "cf-claude",
|
||||
Model: "cloudflare/anthropic/claude-sonnet-4.6",
|
||||
APIKey: "sk-ant-key",
|
||||
CfToken: "cf-test-token",
|
||||
APIBase: "https://gateway.ai.cloudflare.com/v1/acct/gw/compat",
|
||||
}
|
||||
|
||||
provider, modelID, err := CreateProviderFromConfig(cfg)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateProviderFromConfig() error = %v", err)
|
||||
}
|
||||
if provider == nil {
|
||||
t.Fatal("expected non-nil provider")
|
||||
}
|
||||
if _, ok := provider.(*CloudflareProvider); !ok {
|
||||
t.Fatalf("expected *CloudflareProvider, got %T", provider)
|
||||
}
|
||||
if modelID != "anthropic/claude-sonnet-4.6" {
|
||||
t.Errorf("modelID = %q, want %q", modelID, "anthropic/claude-sonnet-4.6")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateProviderFromConfig_Cloudflare_WorkersAI(t *testing.T) {
|
||||
cfg := &config.ModelConfig{
|
||||
ModelName: "cf-gpt-oss-120b",
|
||||
Model: "cloudflare/workers-ai/@cf/openai/gpt-oss-120b",
|
||||
CfToken: "cf-test-token",
|
||||
APIBase: "https://gateway.ai.cloudflare.com/v1/acct/gw/compat",
|
||||
}
|
||||
|
||||
provider, modelID, err := CreateProviderFromConfig(cfg)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateProviderFromConfig() error = %v", err)
|
||||
}
|
||||
if provider == nil {
|
||||
t.Fatal("expected non-nil provider")
|
||||
}
|
||||
if modelID != "workers-ai/@cf/openai/gpt-oss-120b" {
|
||||
t.Errorf("modelID = %q, want %q", modelID, "workers-ai/@cf/openai/gpt-oss-120b")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateProviderFromConfig_Cloudflare_MissingAPIBase(t *testing.T) {
|
||||
cfg := &config.ModelConfig{
|
||||
ModelName: "cf-no-base",
|
||||
Model: "cloudflare/openai/gpt-5.2",
|
||||
CfToken: "cf-test-token",
|
||||
}
|
||||
|
||||
_, _, err := CreateProviderFromConfig(cfg)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for missing api_base")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "api_base is required") {
|
||||
t.Errorf("error = %q, expected to contain 'api_base is required'", err.Error())
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateProviderFromConfig_Cloudflare_MissingAuth(t *testing.T) {
|
||||
cfg := &config.ModelConfig{
|
||||
ModelName: "cf-no-auth",
|
||||
Model: "cloudflare/openai/gpt-5.2",
|
||||
APIBase: "https://gateway.ai.cloudflare.com/v1/acct/gw/compat",
|
||||
}
|
||||
|
||||
_, _, err := CreateProviderFromConfig(cfg)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for missing cf_token and api_key")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "cf_token or api_key is required") {
|
||||
t.Errorf("error = %q, expected to contain 'cf_token or api_key is required'", err.Error())
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateProviderFromConfig_UnknownProtocol(t *testing.T) {
|
||||
cfg := &config.ModelConfig{
|
||||
ModelName: "test-unknown",
|
||||
|
|
|
|||
|
|
@ -31,6 +31,7 @@ type Provider struct {
|
|||
apiKey string
|
||||
apiBase string
|
||||
maxTokensField string // Field name for max tokens (e.g., "max_completion_tokens" for o1/glm models)
|
||||
extraHeaders map[string]string
|
||||
httpClient *http.Client
|
||||
}
|
||||
|
||||
|
|
@ -52,6 +53,13 @@ func WithRequestTimeout(timeout time.Duration) Option {
|
|||
}
|
||||
}
|
||||
|
||||
// WithExtraHeaders adds custom HTTP headers to every request.
|
||||
func WithExtraHeaders(headers map[string]string) Option {
|
||||
return func(p *Provider) {
|
||||
p.extraHeaders = headers
|
||||
}
|
||||
}
|
||||
|
||||
func NewProvider(apiKey, apiBase, proxy string, opts ...Option) *Provider {
|
||||
client := &http.Client{
|
||||
Timeout: defaultRequestTimeout,
|
||||
|
|
@ -176,6 +184,11 @@ func (p *Provider) Chat(
|
|||
req.Header.Set("Authorization", "Bearer "+p.apiKey)
|
||||
}
|
||||
|
||||
// Apply extra headers
|
||||
for k, v := range p.extraHeaders {
|
||||
req.Header.Set(k, v)
|
||||
}
|
||||
|
||||
resp, err := p.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to send request: %w", err)
|
||||
|
|
|
|||
|
|
@ -361,3 +361,63 @@ func TestProvider_FunctionalOptionRequestTimeoutNonPositive(t *testing.T) {
|
|||
t.Fatalf("http timeout = %v, want %v", p.httpClient.Timeout, defaultRequestTimeout)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProvider_FunctionalOptionExtraHeaders(t *testing.T) {
|
||||
headers := map[string]string{
|
||||
"cf-aig-authorization": "Bearer test-token",
|
||||
"X-Custom-Header": "custom-value",
|
||||
}
|
||||
p := NewProvider("key", "https://example.com/v1", "", WithExtraHeaders(headers))
|
||||
if len(p.extraHeaders) != 2 {
|
||||
t.Fatalf("len(extraHeaders) = %d, want 2", len(p.extraHeaders))
|
||||
}
|
||||
if p.extraHeaders["cf-aig-authorization"] != "Bearer test-token" {
|
||||
t.Fatalf("extraHeaders[cf-aig-authorization] = %q, want %q",
|
||||
p.extraHeaders["cf-aig-authorization"], "Bearer test-token")
|
||||
}
|
||||
}
|
||||
|
||||
func TestProvider_ExtraHeaders_SentInRequest(t *testing.T) {
|
||||
var receivedHeaders http.Header
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
receivedHeaders = r.Header.Clone()
|
||||
resp := map[string]any{
|
||||
"choices": []map[string]any{
|
||||
{
|
||||
"message": map[string]any{"content": "ok"},
|
||||
"finish_reason": "stop",
|
||||
},
|
||||
},
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(resp)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
headers := map[string]string{
|
||||
"cf-aig-authorization": "Bearer cf-token-123",
|
||||
}
|
||||
p := NewProvider("api-key", server.URL, "", WithExtraHeaders(headers))
|
||||
|
||||
_, err := p.Chat(
|
||||
t.Context(),
|
||||
[]Message{{Role: "user", Content: "test"}},
|
||||
nil,
|
||||
"openai/gpt-4o",
|
||||
nil,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("Chat() error = %v", err)
|
||||
}
|
||||
|
||||
// Verify extra header is sent
|
||||
if got := receivedHeaders.Get("cf-aig-authorization"); got != "Bearer cf-token-123" {
|
||||
t.Errorf("cf-aig-authorization = %q, want %q", got, "Bearer cf-token-123")
|
||||
}
|
||||
|
||||
// Verify standard Authorization header is also sent
|
||||
if got := receivedHeaders.Get("Authorization"); got != "Bearer api-key" {
|
||||
t.Errorf("Authorization = %q, want %q", got, "Bearer api-key")
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue