Files
spark-mcp/internal/config/config.go
T
tao.chenandClaude f374ebfdb4 uploads,tools,config,main: 将 upload_file 返回路径改为服务器端绝对路径并透传给 spark_submit
之前 upload_file 把文件写到 DataDir/uploads/<filename>,返回相对路径
uploads/hello.txt。MCP 服务器与 agent 不在同一台机器,agent 无法构造服务
器本地路径,因此 spark_submit 无法定位脚本。

改为由新的 internal/uploads 包 mint 一个 32 字符十六进制 file_id,文件落盘为
DataDir/uploads/<file_id>,元数据写入 <file_id>.meta.json(原始文件名只存在
sidecar 里)。upload_file 返回 {file_id, path, name, size, sha256},其中 path
是绝对路径。spark_submit 的 description 明确要求 LLM 直接把 upload_file 返回
的 path 放进 args,不要自己构造路径。

为什么只在描述里约束而不在代码里拒绝非 mint 的绝对路径:集群本地已有路径
(如 /opt/spark/examples/pi.py)是合法的 spark-submit 参数,代码不能替
LLM 拒绝。

测试锁定:
- internal/uploads: Save 往返、AbsPath/Validate 非法路径、Sweep 过期/未过期/
  孤立 sidecar
- internal/mcp/tools: upload_file 新响应字段、spark_submit 透传 mint 路径与
  集群本地路径

刻意未做:S3/HDFS 上传、给 spark_submit 新增 file_id 参数、在代码层面拒绝
集群本地绝对路径。

破坏性变更:upload_file 响应从相对 path 改为绝对 path,并新增 file_id/name/
sha256 字段。该工具尚无外部调用者。

Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-13 19:13:14 +08:00

325 lines
8.2 KiB
Go

// Package config loads server configuration from environment variables.
package config
import (
"bufio"
"fmt"
"log/slog"
"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
GinMode string
AnalyzerDataSkewRatio float64
AnalyzerGCPressureRatio float64
AnalyzerBottleneckShuffleGB float64
UploadTTL time.Duration
}
// 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) {
if err := loadDotenv(".env"); err != nil {
return nil, err
}
cfg := &Config{
ListenAddr: ":8080",
DataDir: "./data",
HTTPClientTimeout: 30 * time.Second,
MaxResponseBytes: 1 << 20,
SparkSubmitTimeout: 60 * time.Second,
UploadTTL: 168 * time.Hour,
LogDir: "./data/logs",
LogLevel: "info",
LogFormat: "text",
GinMode: "release",
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.UploadTTL, err = parseDuration("UPLOAD_TTL", cfg.UploadTTL)
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.GinMode = envString("GIN_MODE", cfg.GinMode)
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
}
if err := validateGinMode(cfg.GinMode); err != nil {
return nil, err
}
if err := validateUploadTTL(cfg.UploadTTL); 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)
}
func validateGinMode(mode string) error {
switch mode {
case "debug", "release", "test":
return nil
}
return fmt.Errorf("config: invalid GIN_MODE %q, want debug/release/test", mode)
}
func validateUploadTTL(d time.Duration) error {
if d <= 0 {
return fmt.Errorf("config: UPLOAD_TTL must be positive")
}
if d > 8760*time.Hour {
return fmt.Errorf("config: UPLOAD_TTL must be at most 8760h (1 year)")
}
return nil
}
// loadDotenv reads KEY=VALUE pairs from path and sets them via os.Setenv only
// when the variable is not already defined. Missing files are ignored.
func loadDotenv(path string) error {
f, err := os.Open(path)
if err != nil {
if os.IsNotExist(err) {
return nil
}
return fmt.Errorf("config: failed to open %s: %w", path, err)
}
defer f.Close()
scanner := bufio.NewScanner(f)
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if line == "" || strings.HasPrefix(line, "#") {
continue
}
eq := strings.Index(line, "=")
if eq < 0 {
slog.Warn("config: .env line without '=', skipping", "line", line)
continue
}
key := strings.TrimSpace(line[:eq])
value := strings.TrimSpace(line[eq+1:])
value = dequote(value)
if key == "" {
slog.Warn("config: .env line with empty key, skipping", "line", line)
continue
}
if os.Getenv(key) == "" {
if err := os.Setenv(key, value); err != nil {
return fmt.Errorf("config: failed to set env %s: %w", key, err)
}
}
}
if err := scanner.Err(); err != nil {
return fmt.Errorf("config: failed to read %s: %w", path, err)
}
return nil
}
func dequote(s string) string {
if len(s) >= 2 {
if (s[0] == '"' && s[len(s)-1] == '"') || (s[0] == '\'' && s[len(s)-1] == '\'') {
return s[1 : len(s)-1]
}
}
return s
}
// 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 UploadTTL=%s "+
"LogDir=%s LogLevel=%s LogFormat=%s AnalyzerDataSkewRatio=%.1f "+
"AnalyzerGCPressureRatio=%.1f AnalyzerBottleneckShuffleGB=%.1f GinMode=%s",
c.ListenAddr,
c.DataDir,
c.SQLitePath,
adminSummary,
agentSummary,
c.HTTPClientTimeout,
c.MaxResponseBytes,
c.SparkSubmitTimeout,
c.UploadTTL,
c.LogDir,
c.LogLevel,
c.LogFormat,
c.AnalyzerDataSkewRatio,
c.AnalyzerGCPressureRatio,
c.AnalyzerBottleneckShuffleGB,
c.GinMode,
)
}