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:
hmes98318 2026-02-27 13:46:18 +08:00
parent 2c8416e658
commit 8d09a2b3e2
11 changed files with 1103 additions and 659 deletions

982
README.md

File diff suppressed because it is too large Load diff

View file

@ -464,6 +464,7 @@ Agent 读取 HEARTBEAT.md
| **神算云** | `shengsuanyun/` | `https://router.shengsuanyun.com/api/v1` | OpenAI | - | | **神算云** | `shengsuanyun/` | `https://router.shengsuanyun.com/api/v1` | OpenAI | - |
| **Antigravity** | `antigravity/` | Google Cloud | 自定义 | 仅 OAuth | | **Antigravity** | `antigravity/` | Google Cloud | 自定义 | 仅 OAuth |
| **GitHub Copilot** | `github-copilot/` | `localhost:4321` | gRPC | - | | **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 GatewayBYOK — 自带密钥)**
通过 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 会自动在它们之间轮询: 为同一个模型名称配置多个端点——PicoClaw 会自动在它们之间轮询:

View file

@ -43,6 +43,12 @@
"model": "openai/gpt-5.2", "model": "openai/gpt-5.2",
"api_key": "sk-key2", "api_key": "sk-key2",
"api_base": "https://api2.example.com/v1" "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": { "channels": {

View file

@ -405,6 +405,9 @@ type ModelConfig struct {
ConnectMode string `json:"connect_mode,omitempty"` // Connection mode: stdio, grpc ConnectMode string `json:"connect_mode,omitempty"` // Connection mode: stdio, grpc
Workspace string `json:"workspace,omitempty"` // Workspace path for CLI-based providers 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 // Optional optimizations
RPM int `json:"rpm,omitempty"` // Requests per minute limit 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") MaxTokensField string `json:"max_tokens_field,omitempty"` // Field name for max tokens (e.g., "max_completion_tokens")

View file

@ -270,6 +270,18 @@ func DefaultConfig() *Config {
APIBase: "http://localhost:8000/v1", APIBase: "http://localhost:8000/v1",
APIKey: "", 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{ Gateway: GatewayConfig{
Host: "127.0.0.1", Host: "127.0.0.1",

View 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 ""
}

View 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)
}
}

View file

@ -53,7 +53,7 @@ func ExtractProtocol(model string) (protocol, modelID string) {
// CreateProviderFromConfig creates a provider based on the ModelConfig. // CreateProviderFromConfig creates a provider based on the ModelConfig.
// It uses the protocol prefix in the Model field to determine which provider to create. // 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. // Returns the provider, the model ID (without protocol prefix), and any error.
func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, error) { func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, error) {
if cfg == nil { if cfg == nil {
@ -139,6 +139,36 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
case "antigravity": case "antigravity":
return NewAntigravityProvider(), modelID, nil 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": case "claude-cli", "claudecli":
workspace := cfg.Workspace workspace := cfg.Workspace
if workspace == "" { if workspace == "" {

View file

@ -64,6 +64,24 @@ func TestExtractProtocol(t *testing.T) {
wantProtocol: "nvidia", wantProtocol: "nvidia",
wantModelID: "meta/llama-3.1-8b", 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 { 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) { func TestCreateProviderFromConfig_UnknownProtocol(t *testing.T) {
cfg := &config.ModelConfig{ cfg := &config.ModelConfig{
ModelName: "test-unknown", ModelName: "test-unknown",

View file

@ -31,6 +31,7 @@ type Provider struct {
apiKey string apiKey string
apiBase string apiBase string
maxTokensField string // Field name for max tokens (e.g., "max_completion_tokens" for o1/glm models) maxTokensField string // Field name for max tokens (e.g., "max_completion_tokens" for o1/glm models)
extraHeaders map[string]string
httpClient *http.Client 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 { func NewProvider(apiKey, apiBase, proxy string, opts ...Option) *Provider {
client := &http.Client{ client := &http.Client{
Timeout: defaultRequestTimeout, Timeout: defaultRequestTimeout,
@ -176,6 +184,11 @@ func (p *Provider) Chat(
req.Header.Set("Authorization", "Bearer "+p.apiKey) 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) resp, err := p.httpClient.Do(req)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to send request: %w", err) return nil, fmt.Errorf("failed to send request: %w", err)

View file

@ -361,3 +361,63 @@ func TestProvider_FunctionalOptionRequestTimeoutNonPositive(t *testing.T) {
t.Fatalf("http timeout = %v, want %v", p.httpClient.Timeout, defaultRequestTimeout) 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")
}
}