Files
spark-mcp/internal/config/config.go
T
tao.chenandClaude a4e2472716 Phase 0+1+2 续: 项目骨架 + 配置 + 日志 + 存储层
Phase 0 (项目骨架):
- main.go 改为最小骨架 (config + logging + gin + /healthz + 优雅退出)
- internal/{config,logging,storage,cluster,...}/ 目录占位

Phase 1 (配置 + 日志):
- internal/config: env 解析 + 必填校验 + token 隐藏的 String()
- internal/logging: slog multi-handler 双输出 (终端 text + 主文件 JSON)
  + StartToolCall per-tool 独立文件 (0600), tools 目录 0700

Phase 2 续 (存储层):
- internal/cluster: 16 字段 Cluster struct (AuthPassword json:"-")
- internal/storage: modernc.org/sqlite 接入, WAL 模式, schema 自动迁移
- ClusterRepo CRUD: Create/Get/List/Update/Delete + ErrNotFound
  + AuthPassword 空字符串 = 保留旧密码 (核心约定)
- 7 个单元测试全绿 (:memory: DB)
- created_at/updated_at 改纳秒精度, 消除 List 测试的 sleep 特殊 case

Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-10 12:12:53 +08:00

226 lines
5.9 KiB
Go

// Package config loads server configuration from environment variables.
package config
import (
"fmt"
"os"
"path/filepath"
"strconv"
"strings"
"time"
)
// Config holds runtime configuration for the Spark MCP server.
type Config struct {
ListenAddr string
DataDir string
SQLitePath string
AdminTokens []string
AgentToken string
HTTPClientTimeout time.Duration
MaxResponseBytes int64
SparkSubmitTimeout time.Duration
LogDir string
LogLevel string
LogFormat string
AnalyzerDataSkewRatio float64
AnalyzerGCPressureRatio float64
AnalyzerBottleneckShuffleGB float64
}
// Load reads configuration from environment variables and returns a populated Config.
// Required variables (ADMIN_TOKENS, AGENT_TOKEN) produce clear errors when missing.
func Load() (*Config, error) {
cfg := &Config{
ListenAddr: ":8080",
DataDir: "./data",
HTTPClientTimeout: 30 * time.Second,
MaxResponseBytes: 1 << 20,
SparkSubmitTimeout: 60 * time.Second,
LogDir: "./data/logs",
LogLevel: "info",
LogFormat: "text",
AnalyzerDataSkewRatio: 3.0,
AnalyzerGCPressureRatio: 0.1,
AnalyzerBottleneckShuffleGB: 50.0,
}
cfg.ListenAddr = envString("LISTEN_ADDR", cfg.ListenAddr)
cfg.DataDir = envString("DATA_DIR", cfg.DataDir)
if err := loadAdminTokens(cfg); err != nil {
return nil, err
}
if err := loadAgentToken(cfg); err != nil {
return nil, err
}
var err error
cfg.HTTPClientTimeout, err = parseDuration("HTTP_CLIENT_TIMEOUT", cfg.HTTPClientTimeout)
if err != nil {
return nil, err
}
cfg.MaxResponseBytes, err = parseInt64("MAX_RESPONSE_BYTES", cfg.MaxResponseBytes)
if err != nil {
return nil, err
}
cfg.SparkSubmitTimeout, err = parseDuration("SPARK_SUBMIT_TIMEOUT", cfg.SparkSubmitTimeout)
if err != nil {
return nil, err
}
cfg.LogDir = envString("LOG_DIR", cfg.LogDir)
cfg.LogLevel = envString("LOG_LEVEL", cfg.LogLevel)
cfg.LogFormat = envString("LOG_FORMAT", cfg.LogFormat)
cfg.AnalyzerDataSkewRatio, err = parseFloat64("ANALYZER_DATA_SKEW_RATIO", cfg.AnalyzerDataSkewRatio)
if err != nil {
return nil, err
}
cfg.AnalyzerGCPressureRatio, err = parseFloat64("ANALYZER_GC_PRESSURE_RATIO", cfg.AnalyzerGCPressureRatio)
if err != nil {
return nil, err
}
cfg.AnalyzerBottleneckShuffleGB, err = parseFloat64("ANALYZER_BOTTLENECK_SHUFFLE_GB", cfg.AnalyzerBottleneckShuffleGB)
if err != nil {
return nil, err
}
if err := validateLogLevel(cfg.LogLevel); err != nil {
return nil, err
}
if err := validateLogFormat(cfg.LogFormat); err != nil {
return nil, err
}
cfg.SQLitePath = filepath.Join(cfg.DataDir, "spark-mcp.db")
return cfg, nil
}
func loadAdminTokens(cfg *Config) error {
v := os.Getenv("ADMIN_TOKENS")
if v == "" {
return fmt.Errorf("config: required environment variable ADMIN_TOKENS is not set")
}
parts := strings.Split(v, ",")
cfg.AdminTokens = make([]string, 0, len(parts))
for _, p := range parts {
p = strings.TrimSpace(p)
if p == "" {
continue
}
cfg.AdminTokens = append(cfg.AdminTokens, p)
}
if len(cfg.AdminTokens) == 0 {
return fmt.Errorf("config: environment variable ADMIN_TOKENS contains no valid tokens")
}
return nil
}
func loadAgentToken(cfg *Config) error {
cfg.AgentToken = os.Getenv("AGENT_TOKEN")
if cfg.AgentToken == "" {
return fmt.Errorf("config: required environment variable AGENT_TOKEN is not set")
}
return nil
}
func envString(key, defaultValue string) string {
if v := os.Getenv(key); v != "" {
return v
}
return defaultValue
}
func parseDuration(key string, defaultValue time.Duration) (time.Duration, error) {
v := os.Getenv(key)
if v == "" {
return defaultValue, nil
}
d, err := time.ParseDuration(v)
if err != nil {
return 0, fmt.Errorf("config: invalid duration for %s: %w", key, err)
}
return d, nil
}
func parseInt64(key string, defaultValue int64) (int64, error) {
v := os.Getenv(key)
if v == "" {
return defaultValue, nil
}
n, err := strconv.ParseInt(v, 10, 64)
if err != nil {
return 0, fmt.Errorf("config: invalid integer for %s: %w", key, err)
}
return n, nil
}
func parseFloat64(key string, defaultValue float64) (float64, error) {
v := os.Getenv(key)
if v == "" {
return defaultValue, nil
}
f, err := strconv.ParseFloat(v, 64)
if err != nil {
return 0, fmt.Errorf("config: invalid float for %s: %w", key, err)
}
return f, nil
}
func validateLogLevel(level string) error {
switch level {
case "debug", "info", "warn", "error":
return nil
}
return fmt.Errorf("config: invalid LOG_LEVEL %q, want debug/info/warn/error", level)
}
func validateLogFormat(format string) error {
switch format {
case "text", "json":
return nil
}
return fmt.Errorf("config: invalid LOG_FORMAT %q, want text/json", format)
}
// String returns a human-readable representation of the configuration.
// Sensitive values (AdminTokens and AgentToken) are summarized, not printed.
func (c *Config) String() string {
adminTotalLen := 0
for _, t := range c.AdminTokens {
adminTotalLen += len(t)
}
adminSummary := fmt.Sprintf("%d token(s), total_len=%d", len(c.AdminTokens), adminTotalLen)
agentSummary := fmt.Sprintf("len=%d", len(c.AgentToken))
return fmt.Sprintf(
"ListenAddr=%s DataDir=%s SQLitePath=%s AdminTokens=%s AgentToken=%s "+
"HTTPClientTimeout=%s MaxResponseBytes=%d SparkSubmitTimeout=%s "+
"LogDir=%s LogLevel=%s LogFormat=%s AnalyzerDataSkewRatio=%.1f "+
"AnalyzerGCPressureRatio=%.1f AnalyzerBottleneckShuffleGB=%.1f",
c.ListenAddr,
c.DataDir,
c.SQLitePath,
adminSummary,
agentSummary,
c.HTTPClientTimeout,
c.MaxResponseBytes,
c.SparkSubmitTimeout,
c.LogDir,
c.LogLevel,
c.LogFormat,
c.AnalyzerDataSkewRatio,
c.AnalyzerGCPressureRatio,
c.AnalyzerBottleneckShuffleGB,
)
}