之前 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>
325 lines
8.2 KiB
Go
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,
|
|
)
|
|
}
|