mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-08-17 11:37:20 +08:00
v0.3.8
This commit is contained in:
@@ -0,0 +1,46 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
legacyruntime "cursor/internal/runtime"
|
||||
)
|
||||
|
||||
func (store *Store) LegacyRuntimeSnapshot(ctx context.Context) (legacyruntime.RuntimeConfigSnapshot, error) {
|
||||
cfg, err := store.Load(ctx)
|
||||
if err != nil {
|
||||
return legacyruntime.RuntimeConfigSnapshot{}, err
|
||||
}
|
||||
|
||||
adapters := make([]legacyruntime.ModelAdapterConfig, 0, len(cfg.ModelAdapters))
|
||||
for _, item := range cfg.ModelAdapters {
|
||||
adapters = append(adapters, legacyruntime.ModelAdapterConfig{
|
||||
ID: item.ID,
|
||||
DisplayName: item.DisplayName,
|
||||
Type: item.Type,
|
||||
BaseURL: item.BaseURL,
|
||||
APIKey: item.APIKey,
|
||||
TooltipData: item.TooltipData,
|
||||
ModelID: item.ModelID,
|
||||
ReasoningEffort: item.ReasoningEffort,
|
||||
OpenAIEndpoint: item.OpenAIEndpoint,
|
||||
OpenAIExtraParamsEnabled: item.OpenAIExtraParamsEnabled,
|
||||
OpenAIExtraParamsJSON: item.OpenAIExtraParamsJSON,
|
||||
CustomHeadersEnabled: item.CustomHeadersEnabled,
|
||||
CustomHeadersJSON: item.CustomHeadersJSON,
|
||||
AnthropicExtraParamsEnabled: item.AnthropicExtraParamsEnabled,
|
||||
AnthropicExtraParamsJSON: item.AnthropicExtraParamsJSON,
|
||||
ContextWindowTokens: item.ContextWindowTokens,
|
||||
MaxCompletionTokens: item.MaxCompletionTokens,
|
||||
AnthropicMaxTokens: item.AnthropicMaxTokens,
|
||||
AnthropicThinkingEffort: item.AnthropicThinkingEffort,
|
||||
ThinkingBudgetTokens: item.ThinkingBudgetTokens,
|
||||
})
|
||||
}
|
||||
|
||||
return legacyruntime.RuntimeConfigSnapshot{
|
||||
ObservabilityLogEnabled: cfg.Log,
|
||||
ProviderStreamIdleTimeout: cfg.ProviderStreamIdleTimeout,
|
||||
ModelAdapters: adapters,
|
||||
}, nil
|
||||
}
|
||||
@@ -0,0 +1,238 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
legacyruntime "cursor/internal/runtime"
|
||||
)
|
||||
|
||||
const configHotReloadMinInterval = 500 * time.Millisecond
|
||||
|
||||
type Manager struct {
|
||||
store *Store
|
||||
current atomic.Pointer[Config]
|
||||
listenersMu sync.RWMutex
|
||||
listeners []func(Config)
|
||||
reloadMu sync.Mutex
|
||||
snapshot fileSnapshot
|
||||
lastReload time.Time
|
||||
reloadError string
|
||||
}
|
||||
|
||||
func NewManager(ctx context.Context, store *Store) (*Manager, error) {
|
||||
if store == nil {
|
||||
return nil, fmt.Errorf("config store is required")
|
||||
}
|
||||
cfg, err := store.Load(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
manager := &Manager{
|
||||
store: store,
|
||||
snapshot: store.snapshot(),
|
||||
}
|
||||
manager.setCurrent(cfg)
|
||||
return manager, nil
|
||||
}
|
||||
|
||||
func (manager *Manager) Current() Config {
|
||||
if manager == nil {
|
||||
return DefaultConfig()
|
||||
}
|
||||
manager.reloadIfChanged(context.Background())
|
||||
return manager.currentConfig()
|
||||
}
|
||||
|
||||
func (manager *Manager) currentConfig() Config {
|
||||
if manager == nil {
|
||||
return DefaultConfig()
|
||||
}
|
||||
if current := manager.current.Load(); current != nil {
|
||||
return *current
|
||||
}
|
||||
return DefaultConfig()
|
||||
}
|
||||
|
||||
func (manager *Manager) Load(ctx context.Context) (Config, error) {
|
||||
if manager == nil {
|
||||
return DefaultConfig(), nil
|
||||
}
|
||||
manager.reloadIfChanged(ctx)
|
||||
return manager.currentConfig(), nil
|
||||
}
|
||||
|
||||
func (manager *Manager) Save(ctx context.Context, cfg Config) (Config, error) {
|
||||
if manager == nil || manager.store == nil {
|
||||
return Config{}, fmt.Errorf("config manager is not initialized")
|
||||
}
|
||||
normalized, err := manager.store.Save(ctx, cfg)
|
||||
if err != nil {
|
||||
return Config{}, err
|
||||
}
|
||||
manager.setCurrent(normalized)
|
||||
manager.reloadMu.Lock()
|
||||
manager.snapshot = manager.store.snapshot()
|
||||
manager.lastReload = time.Now()
|
||||
manager.reloadError = ""
|
||||
manager.reloadMu.Unlock()
|
||||
manager.notify(normalized)
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func (manager *Manager) LastAgentModelHash() string {
|
||||
if manager == nil {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(manager.Current().LastAgentModelHash)
|
||||
}
|
||||
|
||||
func (manager *Manager) SaveLastAgentModelHash(ctx context.Context, value string) error {
|
||||
if manager == nil {
|
||||
return fmt.Errorf("config manager is not initialized")
|
||||
}
|
||||
normalizedValue := strings.TrimSpace(value)
|
||||
current := manager.Current()
|
||||
if strings.TrimSpace(current.LastAgentModelHash) == normalizedValue {
|
||||
return nil
|
||||
}
|
||||
current.LastAgentModelHash = normalizedValue
|
||||
_, err := manager.Save(ctx, current)
|
||||
return err
|
||||
}
|
||||
|
||||
func (manager *Manager) ProviderStreamIdleTimeout(ctx context.Context) time.Duration {
|
||||
if manager == nil {
|
||||
return time.Duration(DefaultProviderStreamIdleTimeoutSeconds) * time.Second
|
||||
}
|
||||
manager.reloadIfChanged(ctx)
|
||||
seconds := normalizeProviderStreamIdleTimeout(manager.currentConfig().ProviderStreamIdleTimeout)
|
||||
return time.Duration(seconds) * time.Second
|
||||
}
|
||||
|
||||
func (manager *Manager) IsObservabilityLogEnabled(ctx context.Context) bool {
|
||||
if manager == nil {
|
||||
return false
|
||||
}
|
||||
manager.reloadIfChanged(ctx)
|
||||
return manager.currentConfig().Log
|
||||
}
|
||||
|
||||
func (manager *Manager) Subscribe(listener func(Config)) func() {
|
||||
if manager == nil || listener == nil {
|
||||
return func() {}
|
||||
}
|
||||
manager.listenersMu.Lock()
|
||||
manager.listeners = append(manager.listeners, listener)
|
||||
index := len(manager.listeners) - 1
|
||||
manager.listenersMu.Unlock()
|
||||
return func() {
|
||||
manager.listenersMu.Lock()
|
||||
defer manager.listenersMu.Unlock()
|
||||
if index < 0 || index >= len(manager.listeners) {
|
||||
return
|
||||
}
|
||||
manager.listeners[index] = nil
|
||||
}
|
||||
}
|
||||
|
||||
func (manager *Manager) LegacyRuntimeSnapshot(_ context.Context) (legacyruntime.RuntimeConfigSnapshot, error) {
|
||||
cfg := manager.Current()
|
||||
adapters := make([]legacyruntime.ModelAdapterConfig, 0, len(cfg.ModelAdapters))
|
||||
for _, item := range cfg.ModelAdapters {
|
||||
adapters = append(adapters, legacyruntime.ModelAdapterConfig{
|
||||
ID: item.ID,
|
||||
DisplayName: item.DisplayName,
|
||||
Type: item.Type,
|
||||
BaseURL: item.BaseURL,
|
||||
APIKey: item.APIKey,
|
||||
TooltipData: item.TooltipData,
|
||||
ModelID: item.ModelID,
|
||||
ReasoningEffort: item.ReasoningEffort,
|
||||
OpenAIEndpoint: item.OpenAIEndpoint,
|
||||
OpenAIExtraParamsEnabled: item.OpenAIExtraParamsEnabled,
|
||||
OpenAIExtraParamsJSON: item.OpenAIExtraParamsJSON,
|
||||
ContextWindowTokens: item.ContextWindowTokens,
|
||||
MaxCompletionTokens: item.MaxCompletionTokens,
|
||||
AnthropicMaxTokens: item.AnthropicMaxTokens,
|
||||
AnthropicThinkingEffort: item.AnthropicThinkingEffort,
|
||||
ThinkingBudgetTokens: item.ThinkingBudgetTokens,
|
||||
})
|
||||
}
|
||||
return legacyruntime.RuntimeConfigSnapshot{
|
||||
ObservabilityLogEnabled: cfg.Log,
|
||||
ProviderStreamIdleTimeout: cfg.ProviderStreamIdleTimeout,
|
||||
ModelAdapters: adapters,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (manager *Manager) RouteMode(hasUpstreamURL bool) string {
|
||||
if !hasUpstreamURL {
|
||||
return DefaultRoutingMode
|
||||
}
|
||||
if manager == nil {
|
||||
return DefaultRoutingMode
|
||||
}
|
||||
mode := normalizeRoutingMode(manager.Current().Routing.Mode)
|
||||
if mode == "" {
|
||||
return DefaultRoutingMode
|
||||
}
|
||||
return mode
|
||||
}
|
||||
|
||||
func (manager *Manager) setCurrent(cfg Config) {
|
||||
next := cfg
|
||||
manager.current.Store(&next)
|
||||
}
|
||||
|
||||
func (manager *Manager) reloadIfChanged(ctx context.Context) {
|
||||
if manager == nil || manager.store == nil {
|
||||
return
|
||||
}
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
now := time.Now()
|
||||
manager.reloadMu.Lock()
|
||||
if !manager.lastReload.IsZero() && now.Sub(manager.lastReload) < configHotReloadMinInterval {
|
||||
manager.reloadMu.Unlock()
|
||||
return
|
||||
}
|
||||
manager.lastReload = now
|
||||
nextSnapshot := manager.store.snapshot()
|
||||
if nextSnapshot == manager.snapshot {
|
||||
manager.reloadMu.Unlock()
|
||||
return
|
||||
}
|
||||
cfg, err := manager.store.Load(ctx)
|
||||
if err != nil {
|
||||
errText := err.Error()
|
||||
if errText != manager.reloadError {
|
||||
log.Printf("config hot reload skipped path=%s error=%v", manager.store.Path(), err)
|
||||
manager.reloadError = errText
|
||||
}
|
||||
manager.reloadMu.Unlock()
|
||||
return
|
||||
}
|
||||
manager.snapshot = nextSnapshot
|
||||
manager.reloadError = ""
|
||||
manager.setCurrent(cfg)
|
||||
manager.reloadMu.Unlock()
|
||||
manager.notify(cfg)
|
||||
}
|
||||
|
||||
func (manager *Manager) notify(cfg Config) {
|
||||
manager.listenersMu.RLock()
|
||||
listeners := append([]func(Config){}, manager.listeners...)
|
||||
manager.listenersMu.RUnlock()
|
||||
for _, listener := range listeners {
|
||||
if listener != nil {
|
||||
listener(cfg)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,86 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
|
||||
"cursor/internal/modelchannel"
|
||||
legacyruntime "cursor/internal/runtime"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultChannelTimeoutMS = int((2 * 60 * 60) * 1000)
|
||||
defaultChannelContextWindowTokens = 200_000
|
||||
defaultChannelMaxTokens = 65_536
|
||||
defaultChannelThinkingBudget = 4_096
|
||||
defaultChannelAnthropicEffort = "xhigh"
|
||||
)
|
||||
|
||||
func (manager *Manager) SelectChannelForModel(_ context.Context, modelID string) (*legacyruntime.ResolvedChannel, error) {
|
||||
if manager == nil {
|
||||
return nil, legacyruntime.ErrChannelNotAvailable
|
||||
}
|
||||
adapters, err := NormalizeModelAdapterConfigs(manager.Current().ModelAdapters)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return resolveModelAdapterChannel(adapters, modelID)
|
||||
}
|
||||
|
||||
func resolveModelAdapterChannel(adapters []ModelAdapterConfig, requestedModel string) (*legacyruntime.ResolvedChannel, error) {
|
||||
matchIndex, ok := modelchannel.ResolveAdapterIndex(
|
||||
adapters,
|
||||
requestedModel,
|
||||
func(adapter ModelAdapterConfig) string { return adapter.ID },
|
||||
func(adapter ModelAdapterConfig) string { return adapter.ModelID },
|
||||
func(adapter ModelAdapterConfig) string {
|
||||
return modelchannel.BuildLegacyChannelID(adapter.BaseURL, adapter.ModelID, adapter.APIKey, adapter.DisplayName)
|
||||
},
|
||||
)
|
||||
if !ok {
|
||||
return nil, legacyruntime.ErrChannelNotAvailable
|
||||
}
|
||||
matched := adapters[matchIndex]
|
||||
|
||||
resolved := &legacyruntime.ResolvedChannel{
|
||||
ID: strings.TrimSpace(matched.ID),
|
||||
Name: strings.TrimSpace(matched.DisplayName),
|
||||
GroupName: "local",
|
||||
Code: strings.TrimSpace(matched.ID),
|
||||
Provider: strings.TrimSpace(matched.Type),
|
||||
BaseURL: strings.TrimSpace(matched.BaseURL),
|
||||
APIKey: strings.TrimSpace(matched.APIKey),
|
||||
Model: strings.TrimSpace(matched.ModelID),
|
||||
OpenAIEndpoint: strings.TrimSpace(matched.OpenAIEndpoint),
|
||||
OpenAIExtraParamsEnabled: matched.OpenAIExtraParamsEnabled,
|
||||
OpenAIExtraParamsJSON: strings.TrimSpace(matched.OpenAIExtraParamsJSON),
|
||||
CustomHeadersEnabled: matched.CustomHeadersEnabled,
|
||||
CustomHeadersJSON: strings.TrimSpace(matched.CustomHeadersJSON),
|
||||
AnthropicExtraParamsEnabled: matched.AnthropicExtraParamsEnabled,
|
||||
AnthropicExtraParamsJSON: strings.TrimSpace(matched.AnthropicExtraParamsJSON),
|
||||
TimeoutMS: defaultChannelTimeoutMS,
|
||||
ContextWindowTokens: defaultChannelContextWindowTokens,
|
||||
MaxTokens: defaultChannelMaxTokens,
|
||||
ReasoningEffort: strings.TrimSpace(matched.ReasoningEffort),
|
||||
AnthropicMaxTokens: defaultChannelMaxTokens,
|
||||
AnthropicThinkingEffort: defaultChannelAnthropicEffort,
|
||||
ThinkingEnabled: true,
|
||||
ThinkingBudgetTokens: defaultChannelThinkingBudget,
|
||||
}
|
||||
if matched.ContextWindowTokens > 0 {
|
||||
resolved.ContextWindowTokens = matched.ContextWindowTokens
|
||||
}
|
||||
if matched.MaxCompletionTokens > 0 {
|
||||
resolved.MaxTokens = matched.MaxCompletionTokens
|
||||
}
|
||||
if matched.AnthropicMaxTokens > 0 {
|
||||
resolved.AnthropicMaxTokens = matched.AnthropicMaxTokens
|
||||
}
|
||||
if matched.ThinkingBudgetTokens > 0 {
|
||||
resolved.ThinkingBudgetTokens = matched.ThinkingBudgetTokens
|
||||
}
|
||||
if strings.TrimSpace(matched.AnthropicThinkingEffort) != "" {
|
||||
resolved.AnthropicThinkingEffort = strings.TrimSpace(matched.AnthropicThinkingEffort)
|
||||
}
|
||||
return resolved, nil
|
||||
}
|
||||
@@ -0,0 +1,166 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
type Store struct {
|
||||
path string
|
||||
logsRoot string
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
type fileSnapshot struct {
|
||||
exists bool
|
||||
modTime int64
|
||||
size int64
|
||||
}
|
||||
|
||||
func NewStore(path string, logsRoot string) *Store {
|
||||
return &Store{
|
||||
path: strings.TrimSpace(path),
|
||||
logsRoot: strings.TrimSpace(logsRoot),
|
||||
}
|
||||
}
|
||||
|
||||
func (store *Store) Path() string {
|
||||
if store == nil {
|
||||
return ""
|
||||
}
|
||||
return store.path
|
||||
}
|
||||
|
||||
func (store *Store) LogsRoot() string {
|
||||
if store == nil {
|
||||
return ""
|
||||
}
|
||||
return store.logsRoot
|
||||
}
|
||||
|
||||
func (store *Store) snapshot() fileSnapshot {
|
||||
if store == nil || strings.TrimSpace(store.path) == "" {
|
||||
return fileSnapshot{}
|
||||
}
|
||||
info, err := os.Stat(store.path)
|
||||
if err != nil {
|
||||
return fileSnapshot{}
|
||||
}
|
||||
return fileSnapshot{
|
||||
exists: true,
|
||||
modTime: info.ModTime().UnixNano(),
|
||||
size: info.Size(),
|
||||
}
|
||||
}
|
||||
|
||||
func (store *Store) Load(_ context.Context) (Config, error) {
|
||||
if store == nil || strings.TrimSpace(store.path) == "" {
|
||||
return DefaultConfig(), nil
|
||||
}
|
||||
|
||||
store.mu.Lock()
|
||||
defer store.mu.Unlock()
|
||||
|
||||
data, err := os.ReadFile(store.path)
|
||||
if err != nil {
|
||||
if errors.Is(err, os.ErrNotExist) {
|
||||
defaultConfig := DefaultConfig()
|
||||
if err := store.saveLocked(defaultConfig); err != nil {
|
||||
return DefaultConfig(), err
|
||||
}
|
||||
return defaultConfig, nil
|
||||
}
|
||||
return DefaultConfig(), fmt.Errorf("读取用户配置失败: %w", err)
|
||||
}
|
||||
|
||||
var current Config
|
||||
if err := yaml.Unmarshal(data, ¤t); err != nil {
|
||||
return DefaultConfig(), fmt.Errorf("解析用户配置失败: %w", err)
|
||||
}
|
||||
normalized, err := NormalizeConfig(current)
|
||||
if err != nil {
|
||||
return DefaultConfig(), err
|
||||
}
|
||||
if shouldPersistNormalizedConfig(data, current, normalized) {
|
||||
if err := store.saveLocked(normalized); err != nil {
|
||||
return DefaultConfig(), err
|
||||
}
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func (store *Store) Save(_ context.Context, cfg Config) (Config, error) {
|
||||
if store == nil || strings.TrimSpace(store.path) == "" {
|
||||
return Config{}, errors.New("配置存储未初始化")
|
||||
}
|
||||
|
||||
normalized, err := NormalizeConfig(cfg)
|
||||
if err != nil {
|
||||
return Config{}, err
|
||||
}
|
||||
|
||||
store.mu.Lock()
|
||||
defer store.mu.Unlock()
|
||||
|
||||
if err := store.saveLocked(normalized); err != nil {
|
||||
return Config{}, err
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func (store *Store) saveLocked(normalized Config) error {
|
||||
if err := os.MkdirAll(filepath.Dir(store.path), 0o755); err != nil {
|
||||
return fmt.Errorf("创建用户配置目录失败: %w", err)
|
||||
}
|
||||
|
||||
data, err := yaml.Marshal(normalized)
|
||||
if err != nil {
|
||||
return fmt.Errorf("序列化用户配置失败: %w", err)
|
||||
}
|
||||
|
||||
tempPath := store.path + ".tmp"
|
||||
if err := os.WriteFile(tempPath, data, 0o644); err != nil {
|
||||
return fmt.Errorf("写入临时配置失败: %w", err)
|
||||
}
|
||||
if err := os.Rename(tempPath, store.path); err != nil {
|
||||
return fmt.Errorf("保存用户配置失败: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func shouldPersistNormalizedConfig(raw []byte, current Config, normalized Config) bool {
|
||||
if !yamlHasKey(raw, "backendListenAddr") || !yamlHasKey(raw, "proxyListenAddr") {
|
||||
return true
|
||||
}
|
||||
if current.BackendListenAddr != normalized.BackendListenAddr || current.ProxyListenAddr != normalized.ProxyListenAddr {
|
||||
return true
|
||||
}
|
||||
if current.ProviderStreamIdleTimeout == normalized.ProviderStreamIdleTimeout {
|
||||
return false
|
||||
}
|
||||
return yamlHasKey(raw, "providerStreamIdleTimeout")
|
||||
}
|
||||
|
||||
func yamlHasKey(raw []byte, key string) bool {
|
||||
var root yaml.Node
|
||||
if err := yaml.Unmarshal(raw, &root); err != nil {
|
||||
return false
|
||||
}
|
||||
if len(root.Content) == 0 || root.Content[0].Kind != yaml.MappingNode {
|
||||
return false
|
||||
}
|
||||
mapping := root.Content[0]
|
||||
for index := 0; index+1 < len(mapping.Content); index += 2 {
|
||||
if mapping.Content[index].Value == key {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,293 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"cursor/internal/modelchannel"
|
||||
)
|
||||
|
||||
const (
|
||||
DefaultBackendListenAddr = "127.0.0.1:18090"
|
||||
DefaultProxyListenAddr = "127.0.0.1:18080"
|
||||
DefaultFrontendBaseURL = "http://127.0.0.1"
|
||||
DefaultRoutingMode = "local"
|
||||
DefaultProviderStreamIdleTimeoutSeconds = 240
|
||||
MinProviderStreamIdleTimeoutSeconds = 30
|
||||
)
|
||||
|
||||
type ModelAdapterConfig struct {
|
||||
ID string `json:"id,omitempty" yaml:"-"`
|
||||
DisplayName string `json:"displayName" yaml:"displayName"`
|
||||
Type string `json:"type" yaml:"type"`
|
||||
BaseURL string `json:"baseURL" yaml:"baseURL"`
|
||||
APIKey string `json:"apiKey" yaml:"apiKey"`
|
||||
TooltipData string `json:"tooltipData" yaml:"tooltipData"`
|
||||
ModelID string `json:"modelID" yaml:"modelID"`
|
||||
ReasoningEffort string `json:"reasoningEffort" yaml:"reasoningEffort"`
|
||||
OpenAIEndpoint string `json:"openAIEndpoint" yaml:"openAIEndpoint"`
|
||||
OpenAIExtraParamsEnabled bool `json:"openAIExtraParamsEnabled" yaml:"openAIExtraParamsEnabled"`
|
||||
OpenAIExtraParamsJSON string `json:"openAIExtraParamsJSON" yaml:"openAIExtraParamsJSON"`
|
||||
CustomHeadersEnabled bool `json:"customHeadersEnabled" yaml:"customHeadersEnabled"`
|
||||
CustomHeadersJSON string `json:"customHeadersJSON" yaml:"customHeadersJSON"`
|
||||
AnthropicExtraParamsEnabled bool `json:"anthropicExtraParamsEnabled" yaml:"anthropicExtraParamsEnabled"`
|
||||
AnthropicExtraParamsJSON string `json:"anthropicExtraParamsJSON" yaml:"anthropicExtraParamsJSON"`
|
||||
ContextWindowTokens int `json:"contextWindowTokens" yaml:"contextWindowTokens"`
|
||||
MaxCompletionTokens int `json:"maxCompletionTokens" yaml:"maxCompletionTokens"`
|
||||
AnthropicMaxTokens int `json:"anthropicMaxTokens" yaml:"anthropicMaxTokens"`
|
||||
AnthropicThinkingEffort string `json:"anthropicThinkingEffort,omitempty" yaml:"anthropicThinkingEffort,omitempty"`
|
||||
ThinkingBudgetTokens int `json:"thinkingBudgetTokens" yaml:"thinkingBudgetTokens"`
|
||||
}
|
||||
|
||||
type RoutingConfig struct {
|
||||
Mode string `json:"mode" yaml:"mode"`
|
||||
}
|
||||
|
||||
type HomeMetricsConfig struct {
|
||||
IncludeCacheWriteInHitRate bool `json:"includeCacheWriteInHitRate" yaml:"includeCacheWriteInHitRate"`
|
||||
}
|
||||
|
||||
type Config struct {
|
||||
Log bool `json:"log" yaml:"log"`
|
||||
ProviderStreamIdleTimeout int `json:"providerStreamIdleTimeout" yaml:"providerStreamIdleTimeout"`
|
||||
BackendListenAddr string `json:"backendListenAddr" yaml:"backendListenAddr"`
|
||||
ProxyListenAddr string `json:"proxyListenAddr" yaml:"proxyListenAddr"`
|
||||
ModelAdapters []ModelAdapterConfig `json:"modelAdapters" yaml:"modelAdapters"`
|
||||
Routing RoutingConfig `json:"routing" yaml:"routing"`
|
||||
HomeMetrics HomeMetricsConfig `json:"homeMetrics" yaml:"homeMetrics"`
|
||||
LastAgentModelHash string `json:"lastAgentModelHash" yaml:"lastAgentModelHash"`
|
||||
}
|
||||
|
||||
func DefaultConfig() Config {
|
||||
return Config{
|
||||
Log: false,
|
||||
ProviderStreamIdleTimeout: DefaultProviderStreamIdleTimeoutSeconds,
|
||||
BackendListenAddr: DefaultBackendListenAddr,
|
||||
ProxyListenAddr: DefaultProxyListenAddr,
|
||||
ModelAdapters: []ModelAdapterConfig{},
|
||||
Routing: RoutingConfig{
|
||||
Mode: DefaultRoutingMode,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func NormalizeConfig(input Config) (Config, error) {
|
||||
output := DefaultConfig()
|
||||
output.Log = input.Log
|
||||
output.ProviderStreamIdleTimeout = normalizeProviderStreamIdleTimeout(input.ProviderStreamIdleTimeout)
|
||||
backendListenAddr, err := normalizeListenAddr(input.BackendListenAddr, DefaultBackendListenAddr, "backendListenAddr")
|
||||
if err != nil {
|
||||
return Config{}, err
|
||||
}
|
||||
proxyListenAddr, err := normalizeListenAddr(input.ProxyListenAddr, DefaultProxyListenAddr, "proxyListenAddr")
|
||||
if err != nil {
|
||||
return Config{}, err
|
||||
}
|
||||
output.BackendListenAddr = backendListenAddr
|
||||
output.ProxyListenAddr = proxyListenAddr
|
||||
output.HomeMetrics.IncludeCacheWriteInHitRate = input.HomeMetrics.IncludeCacheWriteInHitRate
|
||||
output.LastAgentModelHash = strings.TrimSpace(input.LastAgentModelHash)
|
||||
output.Routing.Mode = normalizeRoutingMode(input.Routing.Mode)
|
||||
if output.Routing.Mode == "" {
|
||||
output.Routing.Mode = DefaultRoutingMode
|
||||
}
|
||||
adapters, err := NormalizeModelAdapterConfigs(input.ModelAdapters)
|
||||
if err != nil {
|
||||
return Config{}, err
|
||||
}
|
||||
output.ModelAdapters = adapters
|
||||
return output, nil
|
||||
}
|
||||
|
||||
func NormalizeModelAdapterConfigs(input []ModelAdapterConfig) ([]ModelAdapterConfig, error) {
|
||||
if len(input) == 0 {
|
||||
return []ModelAdapterConfig{}, nil
|
||||
}
|
||||
|
||||
normalized := make([]ModelAdapterConfig, 0, len(input))
|
||||
seenChannelIDs := make(map[string]struct{}, len(input))
|
||||
for _, item := range input {
|
||||
baseURL, err := modelchannel.NormalizeBaseURL(item.BaseURL)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
nextType := normalizeModelAdapterType(item.Type)
|
||||
next := ModelAdapterConfig{
|
||||
DisplayName: strings.TrimSpace(item.DisplayName),
|
||||
Type: nextType,
|
||||
BaseURL: baseURL,
|
||||
APIKey: strings.TrimSpace(item.APIKey),
|
||||
TooltipData: strings.TrimSpace(item.TooltipData),
|
||||
ModelID: strings.TrimSpace(item.ModelID),
|
||||
ReasoningEffort: normalizeReasoningEffort(item.ReasoningEffort),
|
||||
OpenAIEndpoint: modelchannel.NormalizeOpenAIEndpoint(item.Type, item.OpenAIEndpoint),
|
||||
ContextWindowTokens: normalizeMaxCompletionTokens(item.ContextWindowTokens),
|
||||
MaxCompletionTokens: normalizeMaxCompletionTokens(item.MaxCompletionTokens),
|
||||
AnthropicMaxTokens: normalizeMaxCompletionTokens(item.AnthropicMaxTokens),
|
||||
ThinkingBudgetTokens: normalizeMaxCompletionTokens(item.ThinkingBudgetTokens),
|
||||
}
|
||||
if next.Type == "openai" {
|
||||
next.OpenAIExtraParamsEnabled = item.OpenAIExtraParamsEnabled
|
||||
next.OpenAIExtraParamsJSON = strings.TrimSpace(item.OpenAIExtraParamsJSON)
|
||||
} else if next.Type == "anthropic" {
|
||||
next.AnthropicThinkingEffort = normalizeAnthropicThinkingEffort(item.AnthropicThinkingEffort)
|
||||
next.AnthropicExtraParamsEnabled = item.AnthropicExtraParamsEnabled
|
||||
next.AnthropicExtraParamsJSON = strings.TrimSpace(item.AnthropicExtraParamsJSON)
|
||||
}
|
||||
next.CustomHeadersEnabled = item.CustomHeadersEnabled
|
||||
next.CustomHeadersJSON = strings.TrimSpace(item.CustomHeadersJSON)
|
||||
switch {
|
||||
case next.DisplayName == "":
|
||||
return nil, errors.New("模型适配器 displayName 不能为空")
|
||||
case next.Type == "":
|
||||
return nil, errors.New("模型适配器 type 仅支持 openai 或 anthropic")
|
||||
case next.APIKey == "":
|
||||
return nil, errors.New("模型适配器 apiKey 不能为空")
|
||||
case next.TooltipData == "":
|
||||
return nil, errors.New("模型适配器 tooltipData 不能为空")
|
||||
case next.ModelID == "":
|
||||
return nil, errors.New("模型适配器 modelID 不能为空")
|
||||
case next.Type == "openai" && next.ReasoningEffort == "":
|
||||
return nil, errors.New("模型适配器 reasoningEffort 仅支持 low、medium、high、xhigh")
|
||||
case next.Type == "openai" && next.OpenAIEndpoint == "":
|
||||
return nil, errors.New("模型适配器 openAIEndpoint 仅支持 /v1/responses 或 /v1/chat/completions")
|
||||
case next.Type == "openai" && next.OpenAIExtraParamsEnabled:
|
||||
if err := validateJSONMap(next.OpenAIExtraParamsJSON, "openAIExtraParamsJSON"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
case next.CustomHeadersEnabled:
|
||||
if err := validateHeadersJSON(next.CustomHeadersJSON); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
case next.Type == "anthropic" && next.AnthropicExtraParamsEnabled:
|
||||
if err := validateJSONMap(next.AnthropicExtraParamsJSON, "anthropicExtraParamsJSON"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
case next.Type == "anthropic" && next.AnthropicThinkingEffort == "":
|
||||
return nil, errors.New("模型适配器 anthropicThinkingEffort 仅支持 low、medium、high、xhigh、max")
|
||||
}
|
||||
next.ID = modelchannel.BuildChannelID(next.BaseURL, next.ModelID, next.APIKey, next.DisplayName, next.OpenAIEndpoint)
|
||||
if _, exists := seenChannelIDs[next.ID]; exists {
|
||||
return nil, errors.New("模型适配器渠道不能重复,请检查 url、modelID、apiKey、displayName、endpoint 组合")
|
||||
}
|
||||
seenChannelIDs[next.ID] = struct{}{}
|
||||
normalized = append(normalized, next)
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func validateJSONMap(value string, fieldName string) error {
|
||||
text := strings.TrimSpace(value)
|
||||
if text == "" {
|
||||
return fmt.Errorf("模型适配器 %s 不能为空", fieldName)
|
||||
}
|
||||
var parsed map[string]any
|
||||
if err := json.Unmarshal([]byte(text), &parsed); err != nil {
|
||||
return fmt.Errorf("模型适配器 %s 必须是合法 JSON 对象", fieldName)
|
||||
}
|
||||
if parsed == nil {
|
||||
return fmt.Errorf("模型适配器 %s 必须是 JSON 对象", fieldName)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateHeadersJSON(value string) error {
|
||||
text := strings.TrimSpace(value)
|
||||
if err := validateJSONMap(text, "customHeadersJSON"); err != nil {
|
||||
return err
|
||||
}
|
||||
var parsed map[string]string
|
||||
if err := json.Unmarshal([]byte(text), &parsed); err != nil {
|
||||
return errors.New("模型适配器 customHeadersJSON 的值必须是字符串")
|
||||
}
|
||||
for key := range parsed {
|
||||
if strings.TrimSpace(key) == "" {
|
||||
return errors.New("模型适配器 customHeadersJSON 的请求头名称不能为空")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func normalizeReasoningEffort(value string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(value)) {
|
||||
case "", "medium":
|
||||
return "medium"
|
||||
case "low", "high", "xhigh":
|
||||
return strings.ToLower(strings.TrimSpace(value))
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeAnthropicThinkingEffort(value string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(value)) {
|
||||
case "", "xhigh":
|
||||
return "xhigh"
|
||||
case "low", "medium", "high", "max":
|
||||
return strings.ToLower(strings.TrimSpace(value))
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeListenAddr(value string, defaultValue string, fieldName string) (string, error) {
|
||||
addr := strings.TrimSpace(value)
|
||||
if addr == "" {
|
||||
addr = defaultValue
|
||||
}
|
||||
host, port, err := net.SplitHostPort(addr)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("%s 必须是 host:port 格式", fieldName)
|
||||
}
|
||||
if strings.TrimSpace(host) == "" {
|
||||
return "", fmt.Errorf("%s host 不能为空", fieldName)
|
||||
}
|
||||
parsedPort, err := strconv.Atoi(port)
|
||||
if err != nil || parsedPort < 1 || parsedPort > 65535 {
|
||||
return "", fmt.Errorf("%s port 必须在 1-65535 之间", fieldName)
|
||||
}
|
||||
return net.JoinHostPort(host, strconv.Itoa(parsedPort)), nil
|
||||
}
|
||||
|
||||
func normalizeProviderStreamIdleTimeout(value int) int {
|
||||
if value <= 0 {
|
||||
return DefaultProviderStreamIdleTimeoutSeconds
|
||||
}
|
||||
if value < MinProviderStreamIdleTimeoutSeconds {
|
||||
return MinProviderStreamIdleTimeoutSeconds
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
func normalizeMaxCompletionTokens(value int) int {
|
||||
if value <= 0 {
|
||||
return 0
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
func normalizeModelAdapterType(value string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(value)) {
|
||||
case "openai":
|
||||
return "openai"
|
||||
case "anthropic":
|
||||
return "anthropic"
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeRoutingMode(value string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(value)) {
|
||||
case "", "local":
|
||||
return "local"
|
||||
case "upstream":
|
||||
return "upstream"
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user