iop/packages/go/config/load.go

163 lines
5.1 KiB
Go

package config
import (
"fmt"
"strings"
"github.com/spf13/viper"
)
func Load(cfgFile string) (*NodeConfig, error) {
v := viper.New()
v.SetConfigFile(cfgFile)
setDefaults(v)
if err := v.ReadInConfig(); err != nil {
return nil, err
}
var cfg NodeConfig
if err := v.Unmarshal(&cfg); err != nil {
return nil, err
}
return &cfg, nil
}
func LoadEdge(cfgFile string) (*EdgeConfig, error) {
v := viper.New()
v.SetConfigFile(cfgFile)
setEdgeDefaults(v)
if err := v.ReadInConfig(); err != nil {
return nil, err
}
var cfg EdgeConfig
if err := v.Unmarshal(&cfg); err != nil {
return nil, err
}
if !v.InConfig("console.target") {
if v.InConfig("console.agent") {
cfg.Console.Target = cfg.Console.Agent
} else if v.InConfig("console.model") {
cfg.Console.Target = cfg.Console.Model
}
}
if err := validateOpenAIRoutes(cfg.OpenAI.ModelRoutes); err != nil {
return nil, err
}
if err := validateOpenAIPrincipalTokens(cfg.OpenAI.PrincipalTokens); err != nil {
return nil, err
}
if err := normalizeOpenAIProviderAuth(v, &cfg.OpenAI.ProviderAuth); err != nil {
return nil, err
}
if cfg.LongContextThresholdTokens <= 0 {
return nil, fmt.Errorf("long_context_threshold_tokens must be positive")
}
// Collect all provider IDs from nodes[].providers[] for cross-referencing
// and validate uniqueness within each node.
providerIDs := make(map[string]struct{})
providerByID := make(map[string]NodeProviderConf)
for i := range cfg.Nodes {
kind, err := NormalizeAgentKind(cfg.Nodes[i].AgentKind)
if err != nil {
return nil, fmt.Errorf("nodes[%d] alias=%q: %w", i, cfg.Nodes[i].Alias, err)
}
cfg.Nodes[i].AgentKind = kind
if err := normalizeAdapters(&cfg.Nodes[i].Adapters); err != nil {
name := cfg.Nodes[i].ID
if name == "" {
name = cfg.Nodes[i].Alias
}
return nil, fmt.Errorf("nodes[%d] %q adapters: %w", i, name, err)
}
// Validate unique provider IDs within this node and across all nodes.
seenProviderIDs := make(map[string]struct{}, len(cfg.Nodes[i].Providers))
for j, p := range cfg.Nodes[i].Providers {
if err := p.Validate(); err != nil {
return nil, fmt.Errorf("nodes[%d].providers[%d]: %w", i, j, err)
}
if _, dup := seenProviderIDs[p.ID]; dup {
return nil, fmt.Errorf("nodes[%d].providers: duplicate provider id %q within node", i, p.ID)
}
seenProviderIDs[p.ID] = struct{}{}
if _, globalDup := providerIDs[p.ID]; globalDup {
return nil, fmt.Errorf("nodes[%d].providers[%d]: duplicate provider id %q across nodes (global uniqueness violation)", i, j, p.ID)
}
providerIDs[p.ID] = struct{}{}
providerByID[p.ID] = p
}
}
for i := range cfg.Nodes {
if err := CheckProviderLegacyConflict(i, cfg.Nodes[i].Alias, &cfg.Nodes[i]); err != nil {
return nil, err
}
}
// Build the provider->servedModels index so Validate can check membership.
serveModels := buildProviderServedModelsIndex(cfg.Nodes)
// Validate models[].providers reference valid provider IDs.
seenModelIDs := make(map[string]struct{}, len(cfg.Models))
for i, m := range cfg.Models {
id := strings.TrimSpace(m.ID)
if id == "" {
return nil, fmt.Errorf("models[%d]: id must not be empty", i)
}
if _, dup := seenModelIDs[id]; dup {
return nil, fmt.Errorf("models: duplicate model id %q", id)
}
seenModelIDs[id] = struct{}{}
if err := m.Validate(providerIDs, serveModels); err != nil {
return nil, fmt.Errorf("models[%d]: %w", i, err)
}
if err := validateProviderLongContextBudget(m, providerByID); err != nil {
return nil, fmt.Errorf("models[%d]: %w", i, err)
}
}
return &cfg, nil
}
func setDefaults(v *viper.Viper) {
v.SetDefault("transport.edge_addr", "localhost:9090")
v.SetDefault("reconnect.interval_sec", 10)
v.SetDefault("reconnect.max_attempts", 10)
v.SetDefault("logging.level", "info")
v.SetDefault("metrics.port", 9091)
}
func setEdgeDefaults(v *viper.Viper) {
v.SetDefault("server.listen", "0.0.0.0:9090")
v.SetDefault("bootstrap.listen", "0.0.0.0:18080")
v.SetDefault("bootstrap.artifact_dir", "artifacts")
v.SetDefault("openai.enabled", false)
v.SetDefault("openai.listen", "0.0.0.0:18081")
v.SetDefault("openai.bearer_token", "")
v.SetDefault("openai.adapter", "ollama")
v.SetDefault("openai.session_id", "openai")
v.SetDefault("openai.timeout_sec", 120)
v.SetDefault("openai.strict_output", true)
v.SetDefault("openai.strict_stream_buffer", false)
v.SetDefault("a2a.enabled", false)
v.SetDefault("a2a.listen", "0.0.0.0:8081")
v.SetDefault("a2a.path", "/a2a")
v.SetDefault("a2a.adapter", "cli")
v.SetDefault("a2a.session_id", "a2a")
v.SetDefault("a2a.timeout_sec", 120)
v.SetDefault("logging.level", "info")
v.SetDefault("metrics.port", 19092)
v.SetDefault("tls.enabled", false)
v.SetDefault("console.adapter", "cli")
v.SetDefault("console.target", "claude")
v.SetDefault("console.session_id", "default")
v.SetDefault("console.background", false)
v.SetDefault("console.timeout_sec", 120)
v.SetDefault("control_plane.enabled", false)
v.SetDefault("control_plane.wire_addr", "")
v.SetDefault("control_plane.reconnect_interval_sec", 5)
v.SetDefault("refresh.enabled", false)
v.SetDefault("refresh.listen", "127.0.0.1:19093")
v.SetDefault("long_context_threshold_tokens", 100000)
}