package api import ( "fmt" "math" "os" "path/filepath" "runtime" "time" "bat-api/internal/configtemplate" "github.com/pelletier/go-toml/v2" ) // ConfigTOMLTemplate is the canonical template generated by bat-api. Rust // bat consumes the same file and ignores the [api] section. var ConfigTOMLTemplate = configtemplate.Template // LoadConfigFromCurrentExe loads config.toml next to the running binary. func LoadConfigFromCurrentExe(cfg *Config) error { exe, err := os.Executable() if err != nil { return fmt.Errorf("locate current executable: %w", err) } return LoadConfigFromBinaryDir(filepath.Dir(exe), cfg) } // LoadConfigFromBinaryDir loads only config.toml from binaryDir. Missing // config.toml causes config.toml.example to be created when possible; the // example is never parsed as the active configuration. func LoadConfigFromBinaryDir(binaryDir string, cfg *Config) error { if cfg == nil { return fmt.Errorf("config must not be nil") } configPath := filepath.Join(binaryDir, ConfigFileName) examplePath := filepath.Join(binaryDir, ConfigExampleName) info, err := os.Lstat(configPath) if err != nil { if !os.IsNotExist(err) { return fmt.Errorf("stat %s: %w", configPath, err) } ensureConfigExample(examplePath) return nil } if info.Mode()&os.ModeSymlink != 0 || !info.Mode().IsRegular() { return fmt.Errorf("%s must be a regular file", configPath) } if runtime.GOOS != "windows" && info.Mode().Perm()&0o077 != 0 { return fmt.Errorf("%s permissions must be 0600 or stricter (current %03o)", configPath, info.Mode().Perm()) } data, err := os.ReadFile(configPath) if err != nil { return fmt.Errorf("read %s: %w", configPath, err) } if err := applyAPITOML(data, cfg); err != nil { return fmt.Errorf("parse %s: %w", configPath, err) } return nil } func ensureConfigExample(path string) { if _, err := os.Lstat(path); err == nil || !os.IsNotExist(err) { return } file, err := os.OpenFile(path, os.O_WRONLY|os.O_CREATE|os.O_EXCL, 0o600) if err != nil { return } if _, err := file.WriteString(ConfigTOMLTemplate); err != nil { _ = file.Close() return } _ = file.Close() } func applyAPITOML(data []byte, cfg *Config) error { var document map[string]any if err := toml.Unmarshal(data, &document); err != nil { return err } rawAPI, exists := document["api"] if !exists { return nil } apiValues, ok := rawAPI.(map[string]any) if !ok { return fmt.Errorf("[api] must be a table") } for key, value := range apiValues { if err := applyAPIValue(cfg, key, value); err != nil { return err } } return nil } func applyAPIValue(cfg *Config, key string, raw any) error { field := fmt.Sprintf("[api].%s", key) switch key { case "listen": value, err := tomlString(raw, field) if err != nil { return err } cfg.Listen = value case "public_base_url": value, err := tomlString(raw, field) if err != nil { return err } cfg.PublicBaseURL = value case "state_dir": value, err := tomlString(raw, field) if err != nil { return err } cfg.StateDir = value case "socket_path": value, err := tomlString(raw, field) if err != nil { return err } cfg.SocketPath = value case "resource_root": value, err := tomlString(raw, field) if err != nil { return err } cfg.ResourceRoot = value case "server_info_file": value, err := tomlString(raw, field) if err != nil { return err } cfg.ServerInfoFile = value case "require_indexed": value, err := tomlBool(raw, field) if err != nil { return err } cfg.RequireIndexed = value case "verify_size": value, err := tomlBool(raw, field) if err != nil { return err } cfg.VerifySize = value case "rpc_timeout": value, err := tomlDuration(raw, field) if err != nil { return err } cfg.RPCTimeout = value case "refresh_interval": value, err := tomlDuration(raw, field) if err != nil { return err } cfg.RefreshInterval = value case "auth_token": value, err := tomlString(raw, field) if err != nil { return err } cfg.AuthToken = value case "auth_query_param": value, err := tomlString(raw, field) if err != nil { return err } cfg.AuthQueryParam = value case "auth_exempt_paths": value, err := tomlStringArray(raw, field) if err != nil { return err } cfg.AuthExemptPaths = value case "trust_proxy_headers": value, err := tomlBool(raw, field) if err != nil { return err } cfg.TrustProxyHeaders = value case "access_log": value, err := tomlBool(raw, field) if err != nil { return err } cfg.AccessLog = value case "rate_limit_rps": value, err := tomlNonNegativeNumber(raw, field) if err != nil { return err } cfg.RateLimitRPS = value case "rate_limit_burst": value, err := tomlNonNegativeInt(raw, field) if err != nil { return err } cfg.RateLimitBurst = value case "max_resource_page_limit": value, err := tomlNonNegativeInt(raw, field) if err != nil { return err } cfg.MaxResourcePageLimit = value case "database_url": value, err := tomlString(raw, field) if err != nil { return err } cfg.DatabaseURL = value case "database_password": value, err := tomlString(raw, field) if err != nil { return err } cfg.DatabasePassword = value case "redis_url": value, err := tomlString(raw, field) if err != nil { return err } cfg.RedisURL = value case "redis_password": value, err := tomlString(raw, field) if err != nil { return err } cfg.RedisPassword = value default: return fmt.Errorf("unsupported [api] key %q", key) } return nil } func tomlString(raw any, field string) (string, error) { value, ok := raw.(string) if !ok { return "", fmt.Errorf("%s must be a string (got %T)", field, raw) } return value, nil } func tomlBool(raw any, field string) (bool, error) { value, ok := raw.(bool) if !ok { return false, fmt.Errorf("%s must be a boolean (got %T)", field, raw) } return value, nil } func tomlDuration(raw any, field string) (time.Duration, error) { value, err := tomlString(raw, field) if err != nil { return 0, err } duration, err := time.ParseDuration(value) if err != nil || duration < 0 { return 0, fmt.Errorf("%s must be a non-negative duration", field) } return duration, nil } func tomlNonNegativeNumber(raw any, field string) (float64, error) { var value float64 switch typed := raw.(type) { case int64: value = float64(typed) case float64: value = typed default: return 0, fmt.Errorf("%s must be a non-negative number (got %T)", field, raw) } if value < 0 || math.IsNaN(value) || math.IsInf(value, 0) { return 0, fmt.Errorf("%s must be a non-negative number", field) } return value, nil } func tomlNonNegativeInt(raw any, field string) (int, error) { value, ok := raw.(int64) if !ok || value < 0 || int64(int(value)) != value { return 0, fmt.Errorf("%s must be a non-negative integer (got %T)", field, raw) } return int(value), nil } func tomlStringArray(raw any, field string) ([]string, error) { values, ok := raw.([]any) if !ok { return nil, fmt.Errorf("%s must be an array of strings (got %T)", field, raw) } items := make([]string, len(values)) for index, rawItem := range values { item, ok := rawItem.(string) if !ok { return nil, fmt.Errorf("%s[%d] must be a string (got %T)", field, index, rawItem) } items[index] = item } return items, nil }