diff --git a/.dockerignore b/.dockerignore new file mode 100644 index 0000000..ef98500 --- /dev/null +++ b/.dockerignore @@ -0,0 +1,12 @@ +.git +.gitignore +*.md +.env +.env.* +**/*_test.go +test/ +.idea/ +.vscode/ +Makefile +docker-compose.yml +.DS_Store diff --git a/.env.example b/.env.example new file mode 100644 index 0000000..91ca60d --- /dev/null +++ b/.env.example @@ -0,0 +1,4 @@ +OPS_MCP_ENV=prod +OPS_MCP_PORT=8080 +DB_PASSWORD=change_me_in_production +JWT_SECRET=change_me_in_production diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..2ef4cd8 --- /dev/null +++ b/.gitignore @@ -0,0 +1,19 @@ +# Build +/server +/ops-mcp +*.exe + +# IDE +.idea/ +.vscode/ +*.swp +*.swo + +# OS +.DS_Store +Thumbs.db + +# Env +.env +.env.* +!.env.example diff --git a/.golangci.yml b/.golangci.yml new file mode 100644 index 0000000..16b2dfd --- /dev/null +++ b/.golangci.yml @@ -0,0 +1,21 @@ +run: + timeout: 3m + tests: true + +linters: + enable: + - errcheck + - gosimple + - govet + - ineffassign + - staticcheck + - unused + - gofmt + - goimports + - misspell + - unconvert + +issues: + exclude-use-default: false + max-issues-per-linter: 0 + max-same-issues: 0 diff --git a/Dockerfile b/Dockerfile new file mode 100644 index 0000000..d9a23f5 --- /dev/null +++ b/Dockerfile @@ -0,0 +1,20 @@ +FROM golang:1.24-alpine AS builder + +WORKDIR /app + +COPY go.mod go.sum ./ +RUN go mod download + +COPY . . +RUN CGO_ENABLED=0 GOOS=linux go build -ldflags="-s -w" -o /server ./cmd/server + +FROM scratch + +COPY --from=builder /server /server +COPY --from=builder /etc/ssl/certs/ca-certificates.crt /etc/ssl/certs/ +COPY config/ /config/ + +EXPOSE 8080 +USER 65534 + +ENTRYPOINT ["/server"] diff --git a/Makefile b/Makefile new file mode 100644 index 0000000..271b566 --- /dev/null +++ b/Makefile @@ -0,0 +1,54 @@ +.PHONY: build run test lint image clean + +APP_NAME := ops-mcp +IMAGE_NAME := ops-mcp +GO := go +GOFLAGS := -ldflags="-s -w" + +# ---- Build ---- + +build: + $(GO) build $(GOFLAGS) -o $(APP_NAME) ./cmd/server + +run: build + ./$(APP_NAME) -env=dev + +# ---- Test ---- + +test: + $(GO) test -race -v ./... + +test-integration: + $(GO) test -race -v ./test/... + +# ---- Lint ---- + +lint: + golangci-lint run ./... + +# ---- Docker Compose ---- + +up: + docker compose up -d --build + +up-dev: + docker compose --profile dev up -d --build + +down: + docker compose down + +logs: + docker compose logs -f + +# ---- Docker ---- + +image: + docker build -t $(IMAGE_NAME):latest . + +# ---- Utils ---- + +clean: + rm -f $(APP_NAME) + +tidy: + $(GO) mod tidy diff --git a/README.md b/README.md index 088d425..966d35d 100644 --- a/README.md +++ b/README.md @@ -1 +1,80 @@ -ops-mcp server \ No newline at end of file +# ops-mcp + +生产环境 OPS MCP Server —— 安全的数据库查询 + Docker 容器日志工具。 + +Go 编译为 ~15MB 单二进制,scratch 镜像部署。 + +## 快速开始 + +```bash +make build && ./ops-mcp -env=dev +``` + +## Claude Desktop 配置 + +```json +{ + "mcpServers": { + "ops-mcp": { + "command": "/path/to/ops-mcp", + "args": ["-env=dev"] + } + } +} +``` + +## Docker Compose 部署(推荐) + +```bash +# 1. 配置密钥 +cp .env.example .env +# 编辑 .env,填入真实的 DB_PASSWORD 和 JWT_SECRET + +# 2. 启动(仅 MCP 服务) +make up + +# 3. 启动(带本地 PostgreSQL,用于开发调试) +make up-dev + +# 4. 查看日志 +make logs + +# 5. 更新部署 +git pull && make up # 重新 build 镜像并重启 +``` + +**docker-compose.yml 做了什么:** + +| 项目 | 说明 | +|---|---| +| `build.context` | 从当前目录 `Dockerfile` 构建 | +| `ports` | 映射 `8080`(SSE 模式) | +| `volumes` | 只读挂载 `/var/run/docker.sock`,挂载 `prod.yaml` | +| `read_only: true` | 容器文件系统只读 + tmpfs `/tmp` | +| `security_opt` | `no-new-privileges` 禁止提权 | +| `cap_drop: [ALL]` | 移除所有 Linux capabilities | +| `healthcheck` | 30s 间隔自检 | +| `profiles: [dev]` | PG 仅 `--profile dev` 时启动 | + +## MCP 工具列表 + +| 工具 | 说明 | +|---|---| +| `health` | 服务健康检查、内存、goroutine、连接状态 | +| `info` | 版本号、Go 版本、已加载工具列表 | +| `db_query` | 参数化只读 SQL 查询(四层安全防护) | +| `db_tables` | 列出数据库所有表 | +| `db_table_info` | 查看表结构 | +| `db_explain` | EXPLAIN 分析查询计划 | +| `docker_ps` | 列出 Docker 容器 | +| `docker_logs` | 获取容器日志(白名单 + 大小限制) | +| `docker_inspect` | 查看容器详情(不泄露 Env 密钥) | + +## 开发 + +```bash +make build # 编译 +make test # 测试 +make lint # 代码检查(需安装 golangci-lint) +make image # 仅构建 Docker 镜像 +``` diff --git a/cmd/server/main.go b/cmd/server/main.go new file mode 100644 index 0000000..93af605 --- /dev/null +++ b/cmd/server/main.go @@ -0,0 +1,81 @@ +package main + +import ( + "flag" + "fmt" + "log/slog" + "os" + + "ops-mcp/internal/server" + "ops-mcp/internal/tool" + + "ops-mcp/config" +) + +func main() { + env := flag.String("env", "dev", "运行环境:dev | staging | prod") + flag.Parse() + + // 结构化日志 + slog.SetDefault(slog.New(slog.NewTextHandler(os.Stderr, &slog.HandlerOptions{ + Level: slog.LevelInfo, + }))) + + if err := run(*env); err != nil { + slog.Error("fatal", "err", err) + os.Exit(1) + } +} + +func run(env string) error { + // 1. 加载配置 + cfg, err := config.Load(env) + if err != nil { + fmt.Fprintf(os.Stderr, "Warning: could not load config file, using env vars only: %v\n", err) + cfg = &config.Config{ + Env: env, + Server: config.ServerConfig{Transport: "stdio", Addr: ":8080"}, + } + } + + slog.Info("config loaded", "env", cfg.Env, "transport", cfg.Server.Transport) + + // 2. 收集工具 + var tools []tool.Tool + + // database 工具 + if len(cfg.Databases) > 0 { + dbConfigs := make([]tool.DatabaseConf, 0, len(cfg.Databases)) + for alias, db := range cfg.Databases { + dbConfigs = append(dbConfigs, tool.DatabaseConf{ + Alias: alias, + Driver: db.Driver, + DSN: db.DSN, + }) + } + tools = append(tools, tool.NewDatabaseTool(dbConfigs)) + } + + // docker 工具 + if cfg.Docker.Enabled { + dockerTool, err := tool.NewDockerTool(tool.DockerOpts{ + ContainerPattern: cfg.Docker.ContainerPattern, + MaxLogBytes: cfg.Docker.MaxLogBytes, + MaxLogLines: cfg.Docker.MaxLogLines, + MaxLogSince: cfg.Docker.MaxLogSince, + }) + if err != nil { + slog.Error("create docker tool", "err", err) + } else { + tools = append(tools, dockerTool) + } + } + + // 3. 映射 server 配置 + srvCfg := &server.Config{} + srvCfg.Server.Transport = cfg.Server.Transport + srvCfg.Server.Addr = cfg.Server.Addr + + // 4. 启动 + return server.Run(srvCfg, tools...) +} diff --git a/config/config.go b/config/config.go new file mode 100644 index 0000000..97375d7 --- /dev/null +++ b/config/config.go @@ -0,0 +1,81 @@ +package config + +import ( + "fmt" + "time" + + "github.com/spf13/viper" +) + +type Config struct { + Env string `mapstructure:"env"` + Server ServerConfig `mapstructure:"server"` + Databases map[string]DatabaseConf `mapstructure:"databases"` + Docker DockerConfig `mapstructure:"docker"` + Auth AuthConfig `mapstructure:"auth"` + RateLimit RateLimitConfig `mapstructure:"rate_limit"` +} + +type ServerConfig struct { + Transport string `mapstructure:"transport"` // stdio | sse + Addr string `mapstructure:"addr"` // :8080 for SSE mode +} + +type DatabaseConf struct { + Driver string `mapstructure:"driver"` // postgres | mysql + DSN string `mapstructure:"dsn"` +} + +type DockerConfig struct { + Enabled bool `mapstructure:"enabled"` + Host string `mapstructure:"host"` // unix:///var/run/docker.sock + ContainerPattern string `mapstructure:"container_pattern"` + MaxLogBytes int64 `mapstructure:"max_log_bytes"` + MaxLogLines int `mapstructure:"max_log_lines"` + MaxLogSince time.Duration `mapstructure:"max_log_since"` +} + +type AuthConfig struct { + Enabled bool `mapstructure:"enabled"` + APIKeys map[string]string `mapstructure:"api_keys"` // key → description + JWTSecret string `mapstructure:"jwt_secret"` +} + +type RateLimitConfig struct { + Enabled bool `mapstructure:"enabled"` + Rate float64 `mapstructure:"rate"` // 每秒允许的请求数 + Burst int `mapstructure:"burst"` // 突发允许数 +} + +func Load(env string) (*Config, error) { + v := viper.New() + + v.SetConfigName(env) // dev.yaml / staging.yaml / prod.yaml + v.SetConfigType("yaml") + v.AddConfigPath("./config") + v.AddConfigPath(".") + v.AddConfigPath("/config") + + v.SetEnvPrefix("OPS_MCP") + v.AutomaticEnv() + + if err := v.ReadInConfig(); err != nil { + return nil, fmt.Errorf("read config: %w", err) + } + + // 设置默认值 + v.SetDefault("server.transport", "stdio") + v.SetDefault("server.addr", ":8080") + v.SetDefault("docker.max_log_bytes", 1024*1024) + v.SetDefault("docker.max_log_lines", 500) + v.SetDefault("docker.max_log_since", "1h") + v.SetDefault("rate_limit.rate", 10.0) + v.SetDefault("rate_limit.burst", 20) + + var cfg Config + if err := v.Unmarshal(&cfg); err != nil { + return nil, fmt.Errorf("unmarshal config: %w", err) + } + + return &cfg, nil +} diff --git a/config/dev.yaml b/config/dev.yaml new file mode 100644 index 0000000..4991589 --- /dev/null +++ b/config/dev.yaml @@ -0,0 +1,23 @@ +env: dev + +server: + transport: stdio + addr: ":8080" + +databases: + local_pg: + driver: postgres + dsn: postgres://postgres:postgres@localhost:5432/ops_dev?sslmode=disable + +docker: + enabled: false + host: unix:///var/run/docker.sock + container_pattern: "^(app-|web-|api-).*" + +auth: + enabled: false + +rate_limit: + enabled: false + rate: 50 + burst: 100 diff --git a/config/prod.yaml b/config/prod.yaml new file mode 100644 index 0000000..6312419 --- /dev/null +++ b/config/prod.yaml @@ -0,0 +1,24 @@ +env: prod + +server: + transport: sse + addr: ":8080" + +databases: + main_db: + driver: postgres + dsn: postgres://ops_readonly:${DB_PASSWORD}@db-prod.internal:5432/ops?sslmode=require + +docker: + enabled: true + host: unix:///var/run/docker.sock + container_pattern: "^(app-|web-|api-|worker-).*" + +auth: + enabled: true + jwt_secret: ${JWT_SECRET} + +rate_limit: + enabled: true + rate: 10 + burst: 20 diff --git a/config/staging.yaml b/config/staging.yaml new file mode 100644 index 0000000..4402173 --- /dev/null +++ b/config/staging.yaml @@ -0,0 +1,24 @@ +env: staging + +server: + transport: sse + addr: ":8080" + +databases: + main_db: + driver: postgres + dsn: postgres://ops_readonly:${DB_PASSWORD}@db-staging.internal:5432/ops?sslmode=require + +docker: + enabled: true + host: unix:///var/run/docker.sock + container_pattern: "^(app-|web-|api-|worker-).*" + +auth: + enabled: true + jwt_secret: ${JWT_SECRET} + +rate_limit: + enabled: true + rate: 20 + burst: 40 diff --git a/docker-compose.yml b/docker-compose.yml new file mode 100644 index 0000000..e0f3025 --- /dev/null +++ b/docker-compose.yml @@ -0,0 +1,57 @@ +version: "3.8" + +services: + ops-mcp: + build: + context: . + dockerfile: Dockerfile + image: ops-mcp:latest + container_name: ops-mcp + restart: unless-stopped + ports: + - "${OPS_MCP_PORT:-8080}:8080" + environment: + - OPS_MCP_ENV=${OPS_MCP_ENV:-prod} + # 数据库密码通过环境变量注入,不写在配置文件里 + - OPS_MCP_DATABASES_MAIN_DB_DSN=postgres://ops_readonly:${DB_PASSWORD}@db:5432/ops?sslmode=disable&default_transaction_read_only=on + - JWT_SECRET=${JWT_SECRET} + volumes: + # 只读挂载 Docker socket,允许查询容器信息但无法执行写操作 + - /var/run/docker.sock:/var/run/docker.sock:ro + # 挂载自定义配置文件(可选) + - ./config/prod.yaml:/config/prod.yaml:ro + # 安全加固 + read_only: true + tmpfs: + - /tmp:size=10M,mode=1777 + security_opt: + - no-new-privileges:true + cap_drop: + - ALL + # 健康检查 + healthcheck: + test: ["CMD", "/server", "-env=prod"] + interval: 30s + timeout: 5s + retries: 3 + start_period: 10s + + # 可选:本地开发用的 PostgreSQL + db: + image: postgres:17-alpine + container_name: ops-mcp-db + restart: unless-stopped + environment: + - POSTGRES_USER=ops_readonly + - POSTGRES_PASSWORD=${DB_PASSWORD:-devpassword} + - POSTGRES_DB=ops + ports: + - "5432:5432" + volumes: + - pgdata:/var/lib/postgresql/data + profiles: + - dev + - full + +volumes: + pgdata: diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..e63bcae --- /dev/null +++ b/go.mod @@ -0,0 +1,43 @@ +module ops-mcp + +go 1.25.0 + +require ( + github.com/go-sql-driver/mysql v1.9.3 + github.com/golang-jwt/jwt/v5 v5.3.0 + github.com/jackc/pgx/v5 v5.7.6 + github.com/mark3labs/mcp-go v0.42.0 + github.com/spf13/viper v1.21.0 + golang.org/x/time v0.14.0 +) + +require ( + filippo.io/edwards25519 v1.1.0 // indirect + github.com/bahlo/generic-list-go v0.2.0 // indirect + github.com/buger/jsonparser v1.1.1 // indirect + github.com/fsnotify/fsnotify v1.9.0 // indirect + github.com/go-viper/mapstructure/v2 v2.4.0 // indirect + github.com/google/go-cmp v0.7.0 // indirect + github.com/google/uuid v1.6.0 // indirect + github.com/invopop/jsonschema v0.13.0 // indirect + github.com/jackc/pgpassfile v1.0.0 // indirect + github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect + github.com/jackc/puddle/v2 v2.2.2 // indirect + github.com/mailru/easyjson v0.7.7 // indirect + github.com/pelletier/go-toml/v2 v2.2.4 // indirect + github.com/rogpeppe/go-internal v1.14.1 // indirect + github.com/sagikazarmark/locafero v0.11.0 // indirect + github.com/sourcegraph/conc v0.3.1-0.20240121214520-5f936abd7ae8 // indirect + github.com/spf13/afero v1.15.0 // indirect + github.com/spf13/cast v1.10.0 // indirect + github.com/spf13/pflag v1.0.10 // indirect + github.com/subosito/gotenv v1.6.0 // indirect + github.com/wk8/go-ordered-map/v2 v2.1.8 // indirect + github.com/yosida95/uritemplate/v3 v3.0.2 // indirect + go.yaml.in/yaml/v3 v3.0.4 // indirect + golang.org/x/crypto v0.37.0 // indirect + golang.org/x/sync v0.20.0 // indirect + golang.org/x/sys v0.45.0 // indirect + golang.org/x/text v0.37.0 // indirect + gopkg.in/yaml.v3 v3.0.1 // indirect +) diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..747d247 --- /dev/null +++ b/go.sum @@ -0,0 +1,89 @@ +filippo.io/edwards25519 v1.1.0 h1:FNf4tywRC1HmFuKW5xopWpigGjJKiJSV0Cqo0cJWDaA= +filippo.io/edwards25519 v1.1.0/go.mod h1:BxyFTGdWcka3PhytdK4V28tE5sGfRvvvRV7EaN4VDT4= +github.com/bahlo/generic-list-go v0.2.0 h1:5sz/EEAK+ls5wF+NeqDpk5+iNdMDXrh3z3nPnH1Wvgk= +github.com/bahlo/generic-list-go v0.2.0/go.mod h1:2KvAjgMlE5NNynlg/5iLrrCCZ2+5xWbdbCW3pNTGyYg= +github.com/buger/jsonparser v1.1.1 h1:2PnMjfWD7wBILjqQbt530v576A/cAbQvEW9gGIpYMUs= +github.com/buger/jsonparser v1.1.1/go.mod h1:6RYKKt7H4d4+iWqouImQ9R2FZql3VbhNgx27UK13J/0= +github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= +github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/frankban/quicktest v1.14.6 h1:7Xjx+VpznH+oBnejlPUj8oUpdxnVs4f8XU8WnHkI4W8= +github.com/frankban/quicktest v1.14.6/go.mod h1:4ptaffx2x8+WTWXmUCuVU6aPUX1/Mz7zb5vbUoiM6w0= +github.com/fsnotify/fsnotify v1.9.0 h1:2Ml+OJNzbYCTzsxtv8vKSFD9PbJjmhYF14k/jKC7S9k= +github.com/fsnotify/fsnotify v1.9.0/go.mod h1:8jBTzvmWwFyi3Pb8djgCCO5IBqzKJ/Jwo8TRcHyHii0= +github.com/go-sql-driver/mysql v1.9.3 h1:U/N249h2WzJ3Ukj8SowVFjdtZKfu9vlLZxjPXV1aweo= +github.com/go-sql-driver/mysql v1.9.3/go.mod h1:qn46aNg1333BRMNU69Lq93t8du/dwxI64Gl8i5p1WMU= +github.com/go-viper/mapstructure/v2 v2.4.0 h1:EBsztssimR/CONLSZZ04E8qAkxNYq4Qp9LvH92wZUgs= +github.com/go-viper/mapstructure/v2 v2.4.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM= +github.com/golang-jwt/jwt/v5 v5.3.0 h1:pv4AsKCKKZuqlgs5sUmn4x8UlGa0kEVt/puTpKx9vvo= +github.com/golang-jwt/jwt/v5 v5.3.0/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE= +github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= +github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= +github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= +github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/invopop/jsonschema v0.13.0 h1:KvpoAJWEjR3uD9Kbm2HWJmqsEaHt8lBUpd0qHcIi21E= +github.com/invopop/jsonschema v0.13.0/go.mod h1:ffZ5Km5SWWRAIN6wbDXItl95euhFz2uON45H2qjYt+0= +github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM= +github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg= +github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo= +github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM= +github.com/jackc/pgx/v5 v5.7.6 h1:rWQc5FwZSPX58r1OQmkuaNicxdmExaEz5A2DO2hUuTk= +github.com/jackc/pgx/v5 v5.7.6/go.mod h1:aruU7o91Tc2q2cFp5h4uP3f6ztExVpyVv88Xl/8Vl8M= +github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo= +github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4= +github.com/josharian/intern v1.0.0/go.mod h1:5DoeVV0s6jJacbCEi61lwdGj/aVlrQvzHFFd8Hwg//Y= +github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE= +github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk= +github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= +github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= +github.com/mailru/easyjson v0.7.7 h1:UGYAvKxe3sBsEDzO8ZeWOSlIQfWFlxbzLZe7hwFURr0= +github.com/mailru/easyjson v0.7.7/go.mod h1:xzfreul335JAWq5oZzymOObrkdz5UnU4kGfJJLY9Nlc= +github.com/mark3labs/mcp-go v0.42.0 h1:gk/8nYJh8t3yroCAOBhNbYsM9TCKvkM13I5t5Hfu6Ls= +github.com/mark3labs/mcp-go v0.42.0/go.mod h1:YnJfOL382MIWDx1kMY+2zsRHU/q78dBg9aFb8W6Thdw= +github.com/pelletier/go-toml/v2 v2.2.4 h1:mye9XuhQ6gvn5h28+VilKrrPoQVanw5PMw/TB0t5Ec4= +github.com/pelletier/go-toml/v2 v2.2.4/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY= +github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= +github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/rogpeppe/go-internal v1.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0tI/otEQ= +github.com/rogpeppe/go-internal v1.14.1/go.mod h1:MaRKkUm5W0goXpeCfT7UZI6fk/L7L7so1lCWt35ZSgc= +github.com/sagikazarmark/locafero v0.11.0 h1:1iurJgmM9G3PA/I+wWYIOw/5SyBtxapeHDcg+AAIFXc= +github.com/sagikazarmark/locafero v0.11.0/go.mod h1:nVIGvgyzw595SUSUE6tvCp3YYTeHs15MvlmU87WwIik= +github.com/sourcegraph/conc v0.3.1-0.20240121214520-5f936abd7ae8 h1:+jumHNA0Wrelhe64i8F6HNlS8pkoyMv5sreGx2Ry5Rw= +github.com/sourcegraph/conc v0.3.1-0.20240121214520-5f936abd7ae8/go.mod h1:3n1Cwaq1E1/1lhQhtRK2ts/ZwZEhjcQeJQ1RuC6Q/8U= +github.com/spf13/afero v1.15.0 h1:b/YBCLWAJdFWJTN9cLhiXXcD7mzKn9Dm86dNnfyQw1I= +github.com/spf13/afero v1.15.0/go.mod h1:NC2ByUVxtQs4b3sIUphxK0NioZnmxgyCrfzeuq8lxMg= +github.com/spf13/cast v1.10.0 h1:h2x0u2shc1QuLHfxi+cTJvs30+ZAHOGRic8uyGTDWxY= +github.com/spf13/cast v1.10.0/go.mod h1:jNfB8QC9IA6ZuY2ZjDp0KtFO2LZZlg4S/7bzP6qqeHo= +github.com/spf13/pflag v1.0.10 h1:4EBh2KAYBwaONj6b2Ye1GiHfwjqyROoF4RwYO+vPwFk= +github.com/spf13/pflag v1.0.10/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg= +github.com/spf13/viper v1.21.0 h1:x5S+0EU27Lbphp4UKm1C+1oQO+rKx36vfCoaVebLFSU= +github.com/spf13/viper v1.21.0/go.mod h1:P0lhsswPGWD/1lZJ9ny3fYnVqxiegrlNrEmgLjbTCAY= +github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= +github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI= +github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= +github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= +github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= +github.com/subosito/gotenv v1.6.0 h1:9NlTDc1FTs4qu0DDq7AEtTPNw6SVm7uBMsUCUjABIf8= +github.com/subosito/gotenv v1.6.0/go.mod h1:Dk4QP5c2W3ibzajGcXpNraDfq2IrhjMIvMSWPKKo0FU= +github.com/wk8/go-ordered-map/v2 v2.1.8 h1:5h/BUHu93oj4gIdvHHHGsScSTMijfx5PeYkE/fJgbpc= +github.com/wk8/go-ordered-map/v2 v2.1.8/go.mod h1:5nJHM5DyteebpVlHnWMV0rPz6Zp7+xBAnxjb1X5vnTw= +github.com/yosida95/uritemplate/v3 v3.0.2 h1:Ed3Oyj9yrmi9087+NczuL5BwkIc4wvTb5zIM+UJPGz4= +github.com/yosida95/uritemplate/v3 v3.0.2/go.mod h1:ILOh0sOhIJR3+L/8afwt/kE++YT040gmv5BQTMR2HP4= +go.yaml.in/yaml/v3 v3.0.4 h1:tfq32ie2Jv2UxXFdLJdh3jXuOzWiL1fo0bu/FbuKpbc= +go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg= +golang.org/x/crypto v0.37.0 h1:kJNSjF/Xp7kU0iB2Z+9viTPMW4EqqsrywMXLJOOsXSE= +golang.org/x/crypto v0.37.0/go.mod h1:vg+k43peMZ0pUMhYmVAWysMK35e6ioLh3wB8ZCAfbVc= +golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4= +golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= +golang.org/x/sys v0.45.0 h1:dO4czNzziLiiXplLQgBCEpCvXQ3dnkn0SdaZSYdQ+FY= +golang.org/x/sys v0.45.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/text v0.37.0 h1:Cqjiwd9eSg8e0QAkyCaQTNHFIIzWtidPahFWR83rTrc= +golang.org/x/text v0.37.0/go.mod h1:a5sjxXGs9hsn/AJVwuElvCAo9v8QYLzvavO5z2PiM38= +golang.org/x/time v0.14.0 h1:MRx4UaLrDotUKUdCIqzPC48t1Y9hANFKIRpNx+Te8PI= +golang.org/x/time v0.14.0/go.mod h1:eL/Oa2bBBK0TkX57Fyni+NgnyQQN4LitPmob2Hjnqw4= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk= +gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q= +gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/internal/middleware/audit.go b/internal/middleware/audit.go new file mode 100644 index 0000000..22c171f --- /dev/null +++ b/internal/middleware/audit.go @@ -0,0 +1,42 @@ +package middleware + +import ( + "log/slog" + "net/http" + "time" +) + +// responseWriter 捕获状态码,实现 http.ResponseWriter。 +type responseWriter struct { + http.ResponseWriter + status int +} + +func (rw *responseWriter) WriteHeader(code int) { + rw.status = code + rw.ResponseWriter.WriteHeader(code) +} + +// Audit 记录每个 HTTP 请求的结构化审计日志。 +func Audit(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + start := time.Now() + rw := &responseWriter{ResponseWriter: w, status: 200} + + next.ServeHTTP(rw, r) + + subject := "anonymous" + if sub, ok := r.Context().Value(KeySubject).(string); ok { + subject = sub + } + + slog.Info("request", + "method", r.Method, + "path", r.URL.Path, + "status", rw.status, + "duration_ms", time.Since(start).Milliseconds(), + "subject", subject, + "remote_addr", r.RemoteAddr, + ) + }) +} diff --git a/internal/middleware/auth.go b/internal/middleware/auth.go new file mode 100644 index 0000000..f014146 --- /dev/null +++ b/internal/middleware/auth.go @@ -0,0 +1,94 @@ +package middleware + +import ( + "context" + "fmt" + "net/http" + "strings" + + "github.com/golang-jwt/jwt/v5" +) + +type contextKey string + +const ( + KeySubject contextKey = "subject" + KeyScopes contextKey = "scopes" +) + +// Auth 验证 Bearer Token(API Key 或 JWT)。 +type Auth struct { + apiKeys map[string]string // key → description + jwtSecret []byte +} + +func NewAuth(apiKeys map[string]string, jwtSecret string) *Auth { + return &Auth{ + apiKeys: apiKeys, + jwtSecret: []byte(jwtSecret), + } +} + +// Middleware 从 HTTP Header 中提取并验证 Token。 +func (a *Auth) Middleware(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + token := extractBearer(r) + if token == "" { + http.Error(w, "missing Bearer token", http.StatusUnauthorized) + return + } + + // 尝试 API Key 验证 + if desc, ok := a.apiKeys[token]; ok { + ctx := context.WithValue(r.Context(), KeySubject, fmt.Sprintf("apikey:%s", desc)) + ctx = context.WithValue(ctx, KeyScopes, []string{"*"}) + next.ServeHTTP(w, r.WithContext(ctx)) + return + } + + // 尝试 JWT 验证 + if a.jwtSecret != nil { + claims, err := a.parseJWT(token) + if err == nil { + ctx := context.WithValue(r.Context(), KeySubject, claims.Subject) + ctx = context.WithValue(ctx, KeyScopes, claims.Scopes) + next.ServeHTTP(w, r.WithContext(ctx)) + return + } + } + + http.Error(w, "invalid token", http.StatusUnauthorized) + }) +} + +type customClaims struct { + jwt.RegisteredClaims + Scopes []string `json:"scopes"` +} + +func (a *Auth) parseJWT(tokenStr string) (*customClaims, error) { + token, err := jwt.ParseWithClaims(tokenStr, &customClaims{}, + func(t *jwt.Token) (any, error) { + return a.jwtSecret, nil + }, + ) + if err != nil { + return nil, err + } + claims, ok := token.Claims.(*customClaims) + if !ok || !token.Valid { + return nil, fmt.Errorf("invalid claims") + } + return claims, nil +} + +func extractBearer(r *http.Request) string { + auth := r.Header.Get("Authorization") + if auth == "" { + return "" + } + if !strings.HasPrefix(auth, "Bearer ") { + return "" + } + return strings.TrimPrefix(auth, "Bearer ") +} diff --git a/internal/middleware/ratelimit.go b/internal/middleware/ratelimit.go new file mode 100644 index 0000000..a3155b1 --- /dev/null +++ b/internal/middleware/ratelimit.go @@ -0,0 +1,53 @@ +package middleware + +import ( + "net/http" + "sync" + + "golang.org/x/time/rate" +) + +// RateLimiter 基于调用方标识的令牌桶限流。 +type RateLimiter struct { + limiters map[string]*rate.Limiter + rate rate.Limit + burst int + mu sync.Mutex +} + +func NewRateLimiter(r rate.Limit, burst int) *RateLimiter { + return &RateLimiter{ + limiters: make(map[string]*rate.Limiter), + rate: r, + burst: burst, + } +} + +func (rl *RateLimiter) Middleware(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + subject := "anonymous" + if sub, ok := r.Context().Value(KeySubject).(string); ok { + subject = sub + } + + limiter := rl.getLimiter(subject) + if !limiter.Allow() { + http.Error(w, "rate limit exceeded", http.StatusTooManyRequests) + return + } + + next.ServeHTTP(w, r) + }) +} + +func (rl *RateLimiter) getLimiter(key string) *rate.Limiter { + rl.mu.Lock() + defer rl.mu.Unlock() + + limiter, ok := rl.limiters[key] + if !ok { + limiter = rate.NewLimiter(rl.rate, rl.burst) + rl.limiters[key] = limiter + } + return limiter +} diff --git a/internal/server/config.go b/internal/server/config.go new file mode 100644 index 0000000..8ce657f --- /dev/null +++ b/internal/server/config.go @@ -0,0 +1,17 @@ +package server + +// Config 是 server 包所需的配置子集,避免循环依赖。 +// 实际值由 main.go 从 config.Config 映射过来。 +type Config struct { + Server struct { + Transport string + Addr string + } + Docker struct { + Enabled bool + } + Databases map[string]struct { + Driver string + DSN string + } +} diff --git a/internal/server/registry.go b/internal/server/registry.go new file mode 100644 index 0000000..5a7abcd --- /dev/null +++ b/internal/server/registry.go @@ -0,0 +1,83 @@ +package server + +import ( + "context" + "fmt" + "log/slog" + + "ops-mcp/internal/tool" + + "github.com/mark3labs/mcp-go/server" +) + +// Registry 管理所有工具的生命周期。 +type Registry struct { + tools map[string]tool.Tool +} + +func NewRegistry() *Registry { + return &Registry{tools: make(map[string]tool.Tool)} +} + +// Register 注册一个工具。如果工具名重复则报错。 +func (r *Registry) Register(t tool.Tool) error { + name := t.Name() + if _, exists := r.tools[name]; exists { + return fmt.Errorf("tool %q already registered", name) + } + r.tools[name] = t + slog.Info("tool registered", "name", name, "desc", t.Description()) + return nil +} + +// InitializeAll 调用所有工具的 Initialize,任一失败则终止。 +func (r *Registry) InitializeAll(ctx context.Context) error { + for name, t := range r.tools { + slog.Info("initializing tool", "name", name) + if err := t.Initialize(ctx); err != nil { + return fmt.Errorf("init tool %q: %w", name, err) + } + } + return nil +} + +// RegisterAll 将所有工具注册到 MCP Server。 +func (r *Registry) RegisterAll(mcpServer *server.MCPServer) error { + for name, t := range r.tools { + if err := t.Register(mcpServer); err != nil { + return fmt.Errorf("register tool %q: %w", name, err) + } + slog.Info("tool handlers registered", "name", name) + } + return nil +} + +// ShutdownAll 调用所有工具的 Shutdown,收集所有错误。 +func (r *Registry) ShutdownAll(ctx context.Context) []error { + var errs []error + for name, t := range r.tools { + slog.Info("shutting down tool", "name", name) + if err := t.Shutdown(ctx); err != nil { + errs = append(errs, fmt.Errorf("shutdown tool %q: %w", name, err)) + } + } + return errs +} + +// HealthCheckAll 检查所有工具的连通性,用于 system health 工具。 +func (r *Registry) HealthCheckAll(ctx context.Context) map[string]error { + result := make(map[string]error, len(r.tools)) + for name, t := range r.tools { + result[name] = t.HealthCheck(ctx) + } + return result +} + +// List 返回所有已注册工具的名称。 +func (r *Registry) List() []string { + names := make([]string, 0, len(r.tools)) + for name := range r.tools { + names = append(names, name) + } + return names +} diff --git a/internal/server/server.go b/internal/server/server.go new file mode 100644 index 0000000..616a7d8 --- /dev/null +++ b/internal/server/server.go @@ -0,0 +1,111 @@ +package server + +import ( + "context" + "fmt" + "log/slog" + "net/http" + "os" + "os/signal" + "syscall" + + "ops-mcp/internal/tool" + + "github.com/mark3labs/mcp-go/server" +) + +const version = "0.1.0" + +// Run 启动 MCP Server,处理信号优雅关闭。 +func Run(cfg *Config, tools ...tool.Tool) error { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + // 监听终止信号 + sigCh := make(chan os.Signal, 1) + signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM) + go func() { + sig := <-sigCh + slog.Info("received signal, shutting down", "signal", sig) + cancel() + }() + + // 创建注册中心 + registry := NewRegistry() + for _, t := range tools { + if err := registry.Register(t); err != nil { + return fmt.Errorf("register: %w", err) + } + } + + // 系统工具总是最后注册,确保它能看到所有工具 + sysTool := tool.NewSystemTool(version, registry) + if err := registry.Register(sysTool); err != nil { + return fmt.Errorf("register system tool: %w", err) + } + + // 初始化所有工具 + if err := registry.InitializeAll(ctx); err != nil { + return fmt.Errorf("initialize: %w", err) + } + defer func() { + for _, err := range registry.ShutdownAll(ctx) { + slog.Error("shutdown error", "err", err) + } + }() + + // 创建 MCP Server + mcpServer := server.NewMCPServer( + "ops-mcp", + version, + server.WithLogging(), + ) + + // 注册所有工具处理器 + if err := registry.RegisterAll(mcpServer); err != nil { + return fmt.Errorf("register all: %w", err) + } + + // 选择传输模式 + switch cfg.Server.Transport { + case "stdio": + return runStdio(ctx, mcpServer) + case "sse": + return runSSE(ctx, mcpServer, cfg) + default: + return fmt.Errorf("unknown transport: %s (expect stdio or sse)", cfg.Server.Transport) + } +} + +func runStdio(ctx context.Context, mcpServer *server.MCPServer) error { + slog.Info("starting MCP server", "transport", "stdio") + go func() { + <-ctx.Done() + slog.Info("shutting down stdio server") + }() + return server.ServeStdio(mcpServer) +} + +func runSSE(ctx context.Context, mcpServer *server.MCPServer, cfg *Config) error { + sseServer := server.NewSSEServer(mcpServer) + + addr := cfg.Server.Addr + if addr == "" { + addr = ":8080" + } + + slog.Info("starting MCP server", "transport", "sse", "addr", addr) + + httpServer := &http.Server{Addr: addr, Handler: sseServer} + + go func() { + <-ctx.Done() + slog.Info("shutting down HTTP server") + httpServer.Shutdown(context.Background()) + }() + + if err := httpServer.ListenAndServe(); err != nil && err != http.ErrServerClosed { + return err + } + return nil +} diff --git a/internal/tool/database.go b/internal/tool/database.go new file mode 100644 index 0000000..49135ec --- /dev/null +++ b/internal/tool/database.go @@ -0,0 +1,399 @@ +package tool + +import ( + "context" + "database/sql" + "encoding/json" + "fmt" + "log/slog" + "regexp" + "strings" + "sync" + "time" + + _ "github.com/go-sql-driver/mysql" + "github.com/mark3labs/mcp-go/mcp" + "github.com/mark3labs/mcp-go/server" + + _ "github.com/jackc/pgx/v5/stdlib" +) + +const ( + maxRows = 1000 + maxSQLBytes = 4096 + queryTimeout = 30 * time.Second +) + +// 禁止的 SQL 关键词(大写,用于只读防护的第二层) +var forbiddenKeywords = regexp.MustCompile( + `\b(DROP|TRUNCATE|ALTER|CREATE|INSERT|UPDATE|DELETE|GRANT|REVOKE|REPLACE|LOAD|IMPORT|EXPORT)\b`, +) + +// DatabaseConf 数据库连接配置 +type DatabaseConf struct { + Alias string + Driver string // postgres | mysql + DSN string +} + +// DatabaseTool 提供只读数据库查询工具。 +type DatabaseTool struct { + configs map[string]DatabaseConf + pools map[string]*sql.DB + mu sync.RWMutex +} + +func NewDatabaseTool(configs []DatabaseConf) *DatabaseTool { + cfgMap := make(map[string]DatabaseConf, len(configs)) + for _, c := range configs { + cfgMap[c.Alias] = c + } + return &DatabaseTool{configs: cfgMap, pools: make(map[string]*sql.DB)} +} + +func (d *DatabaseTool) Name() string { return "database" } +func (d *DatabaseTool) Description() string { return "数据库只读查询" } + +func (d *DatabaseTool) Initialize(ctx context.Context) error { + for alias, cfg := range d.configs { + // 强制追加只读参数 + dsn := d.enforceReadOnly(cfg) + pool, err := sql.Open(cfg.Driver, dsn) + if err != nil { + return fmt.Errorf("open %s: %w", alias, err) + } + pool.SetMaxOpenConns(5) + pool.SetMaxIdleConns(2) + pool.SetConnMaxLifetime(5 * time.Minute) + + if err := pool.PingContext(ctx); err != nil { + return fmt.Errorf("ping %s: %w", alias, err) + } + d.pools[alias] = pool + slog.Info("database connected", "alias", alias, "driver", cfg.Driver) + } + return nil +} + +func (d *DatabaseTool) enforceReadOnly(cfg DatabaseConf) string { + switch cfg.Driver { + case "postgres", "pgx": + if !strings.Contains(cfg.DSN, "default_transaction_read_only") { + sep := "?" + if strings.Contains(cfg.DSN, "?") { + sep = "&" + } + return cfg.DSN + sep + "default_transaction_read_only=on" + } + case "mysql": + // mysql 驱动在 DSN 中不支持这个参数,通过连接时设置 + } + return cfg.DSN +} + +func (d *DatabaseTool) Shutdown(_ context.Context) error { + for alias, pool := range d.pools { + if err := pool.Close(); err != nil { + slog.Error("close db pool", "alias", alias, "err", err) + } + } + return nil +} + +func (d *DatabaseTool) HealthCheck(ctx context.Context) error { + for alias, pool := range d.pools { + if err := pool.PingContext(ctx); err != nil { + return fmt.Errorf("%s: %w", alias, err) + } + } + return nil +} + +func (d *DatabaseTool) Register(mcpServer *server.MCPServer) error { + mcpServer.AddTool(mcp.NewTool("db_query", + mcp.WithDescription("执行只读 SQL 查询(参数化)"), + mcp.WithString("database", mcp.Required(), mcp.Description("数据库别名")), + mcp.WithString("sql", mcp.Required(), mcp.Description("SQL 查询语句")), + mcp.WithString("params", mcp.Description("JSON 数组格式的参数,如 [1, 'hello']")), + ), d.handleQuery) + + mcpServer.AddTool(mcp.NewTool("db_tables", + mcp.WithDescription("列出数据库中的所有表"), + mcp.WithString("database", mcp.Required(), mcp.Description("数据库别名")), + mcp.WithString("schema", mcp.Description("schema 名称(PostgreSQL),默认 public")), + ), d.handleTables) + + mcpServer.AddTool(mcp.NewTool("db_table_info", + mcp.WithDescription("查看表结构和索引"), + mcp.WithString("database", mcp.Required(), mcp.Description("数据库别名")), + mcp.WithString("table", mcp.Required(), mcp.Description("表名")), + ), d.handleTableInfo) + + mcpServer.AddTool(mcp.NewTool("db_explain", + mcp.WithDescription("EXPLAIN 分析查询计划"), + mcp.WithString("database", mcp.Required(), mcp.Description("数据库别名")), + mcp.WithString("sql", mcp.Required(), mcp.Description("要分析的 SQL")), + ), d.handleExplain) + + return nil +} + +// --- 安全校验 --- + +func (d *DatabaseTool) validateSQL(sqlStr string) error { + if len(sqlStr) > maxSQLBytes { + return fmt.Errorf("SQL too long: %d bytes (max %d)", len(sqlStr), maxSQLBytes) + } + // 禁止多语句 + if strings.Contains(sqlStr, ";") { + return fmt.Errorf("multiple statements not allowed") + } + // 禁止危险关键词 + upper := strings.ToUpper(sqlStr) + if loc := forbiddenKeywords.FindStringIndex(upper); loc != nil { + return fmt.Errorf("forbidden keyword detected near position %d", loc[0]) + } + return nil +} + +func (d *DatabaseTool) getPool(alias string) (*sql.DB, string, error) { + d.mu.RLock() + defer d.mu.RUnlock() + + pool, ok := d.pools[alias] + if !ok { + return nil, "", fmt.Errorf("unknown database: %s (available: %v)", alias, d.listAliases()) + } + + cfg := d.configs[alias] + return pool, cfg.Driver, nil +} + +func (d *DatabaseTool) listAliases() []string { + aliases := make([]string, 0, len(d.pools)) + for a := range d.pools { + aliases = append(aliases, a) + } + return aliases +} + +func parseParams(raw string) ([]any, error) { + if raw == "" { + return nil, nil + } + var params []any + if err := json.Unmarshal([]byte(raw), ¶ms); err != nil { + return nil, fmt.Errorf("params must be a JSON array: %w", err) + } + return params, nil +} + +// --- Handlers --- + +func (d *DatabaseTool) handleQuery(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { + args := getArgs(req) + alias, _ := args["database"].(string) + sqlStr, _ := args["sql"].(string) + rawParams, _ := args["params"].(string) + + if err := d.validateSQL(sqlStr); err != nil { + return mcp.NewToolResultError(err.Error()), nil + } + + pool, driver, err := d.getPool(alias) + if err != nil { + return mcp.NewToolResultError(err.Error()), nil + } + + params, err := parseParams(rawParams) + if err != nil { + return mcp.NewToolResultError(err.Error()), nil + } + + ctx, cancel := context.WithTimeout(ctx, queryTimeout) + defer cancel() + + // 对 MySQL 连接设置只读 + if driver == "mysql" { + if _, err := pool.ExecContext(ctx, "SET SESSION TRANSACTION READ ONLY"); err == nil { + defer pool.ExecContext(context.Background(), "SET SESSION TRANSACTION READ WRITE") + } + } + + sqlStr = limitQuery(sqlStr, driver) + + rows, err := pool.QueryContext(ctx, sqlStr, params...) + if err != nil { + return mcp.NewToolResultError(fmt.Sprintf("query error: %v", err)), nil + } + defer rows.Close() + + return formatRows(rows, maxRows) +} + +func (d *DatabaseTool) handleTables(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { + args := getArgs(req) + alias, _ := args["database"].(string) + schema, _ := args["schema"].(string) + + pool, driver, err := d.getPool(alias) + if err != nil { + return mcp.NewToolResultError(err.Error()), nil + } + + var query string + switch driver { + case "postgres", "pgx": + if schema == "" { + schema = "public" + } + query = `SELECT table_name FROM information_schema.tables WHERE table_schema = $1 ORDER BY table_name` + case "mysql": + if schema == "" { + schema = "" + } + query = `SELECT table_name FROM information_schema.tables WHERE table_schema = ? ORDER BY table_name` + default: + query = `SELECT table_name FROM information_schema.tables WHERE table_schema = ? ORDER BY table_name` + } + + ctx, cancel := context.WithTimeout(ctx, queryTimeout) + defer cancel() + + var rows *sql.Rows + var err2 error + if schema != "" { + rows, err2 = pool.QueryContext(ctx, query, schema) + } else { + rows, err2 = pool.QueryContext(ctx, query) + } + if err2 != nil { + return mcp.NewToolResultError(err2.Error()), nil + } + defer rows.Close() + + return formatRows(rows, maxRows) +} + +func (d *DatabaseTool) handleTableInfo(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { + args := getArgs(req) + alias, _ := args["database"].(string) + table, _ := args["table"].(string) + + pool, driver, err := d.getPool(alias) + if err != nil { + return mcp.NewToolResultError(err.Error()), nil + } + + // 校验表名只包含安全字符 + if !regexp.MustCompile(`^[a-zA-Z_][a-zA-Z0-9_]*$`).MatchString(table) { + return mcp.NewToolResultError("invalid table name"), nil + } + + var query string + switch driver { + case "postgres", "pgx": + query = `SELECT column_name, data_type, is_nullable, column_default FROM information_schema.columns WHERE table_name = $1 ORDER BY ordinal_position` + case "mysql": + query = `SELECT column_name, data_type, is_nullable, column_default FROM information_schema.columns WHERE table_name = ? ORDER BY ordinal_position` + default: + query = `SELECT column_name, data_type, is_nullable, column_default FROM information_schema.columns WHERE table_name = ? ORDER BY ordinal_position` + } + + ctx, cancel := context.WithTimeout(ctx, queryTimeout) + defer cancel() + + rows, err := pool.QueryContext(ctx, query, table) + if err != nil { + return mcp.NewToolResultError(err.Error()), nil + } + defer rows.Close() + + return formatRows(rows, maxRows) +} + +func (d *DatabaseTool) handleExplain(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { + args := getArgs(req) + alias, _ := args["database"].(string) + sqlStr, _ := args["sql"].(string) + + if err := d.validateSQL(sqlStr); err != nil { + return mcp.NewToolResultError(err.Error()), nil + } + + pool, driver, err := d.getPool(alias) + if err != nil { + return mcp.NewToolResultError(err.Error()), nil + } + + explainPrefix := "EXPLAIN" + if driver == "mysql" { + explainPrefix = "EXPLAIN" + } + + ctx, cancel := context.WithTimeout(ctx, queryTimeout) + defer cancel() + + rows, err := pool.QueryContext(ctx, explainPrefix+" "+sqlStr) + if err != nil { + return mcp.NewToolResultError(err.Error()), nil + } + defer rows.Close() + + return formatRows(rows, maxRows) +} + +// --- Helpers --- + +// limitQuery 追加 LIMIT 防止返回过多行(如果查询中没有 LIMIT)。 +func limitQuery(sqlStr, driver string) string { + upper := strings.ToUpper(sqlStr) + if strings.Contains(upper, "LIMIT") { + return sqlStr + } + return sqlStr + fmt.Sprintf(" LIMIT %d", maxRows) +} + +func formatRows(rows *sql.Rows, limit int) (*mcp.CallToolResult, error) { + columns, err := rows.Columns() + if err != nil { + return mcp.NewToolResultError(err.Error()), nil + } + + var result []map[string]any + count := 0 + for rows.Next() { + if count >= limit { + break + } + values := make([]any, len(columns)) + valuePtrs := make([]any, len(columns)) + for i := range values { + valuePtrs[i] = &values[i] + } + if err := rows.Scan(valuePtrs...); err != nil { + return mcp.NewToolResultError(err.Error()), nil + } + row := make(map[string]any, len(columns)) + for i, col := range columns { + row[col] = formatValue(values[i]) + } + result = append(result, row) + count++ + } + + data, _ := json.MarshalIndent(result, "", " ") + output := fmt.Sprintf("Rows: %d\nColumns: %s\n\n%s", count, strings.Join(columns, ", "), string(data)) + return mcp.NewToolResultText(output), nil +} + +func formatValue(v any) any { + switch val := v.(type) { + case []byte: + return string(val) + case time.Time: + return val.Format(time.RFC3339) + default: + return v + } +} diff --git a/internal/tool/docker.go b/internal/tool/docker.go new file mode 100644 index 0000000..349f77e --- /dev/null +++ b/internal/tool/docker.go @@ -0,0 +1,278 @@ +package tool + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "log/slog" + "os/exec" + "regexp" + "strings" + "time" + + "github.com/mark3labs/mcp-go/mcp" + "github.com/mark3labs/mcp-go/server" +) + +// DockerTool 提供安全的 Docker 容器信息查询工具。 +// 使用 docker CLI 而非 SDK,避免跨平台编译问题,且更轻量。 +type DockerTool struct { + containerRegex *regexp.Regexp + maxLogBytes int64 + maxLogLines int + maxLogSince time.Duration +} + +type DockerOpts struct { + ContainerPattern string + MaxLogBytes int64 + MaxLogLines int + MaxLogSince time.Duration +} + +func NewDockerTool(opts DockerOpts) (*DockerTool, error) { + pattern := opts.ContainerPattern + if pattern == "" { + pattern = ".*" + } + + re, err := regexp.Compile(pattern) + if err != nil { + return nil, fmt.Errorf("compile container pattern: %w", err) + } + + return &DockerTool{ + containerRegex: re, + maxLogBytes: opts.MaxLogBytes, + maxLogLines: opts.MaxLogLines, + maxLogSince: opts.MaxLogSince, + }, nil +} + +func (d *DockerTool) Name() string { return "docker" } +func (d *DockerTool) Description() string { return "Docker 容器日志和安全查询" } + +func (d *DockerTool) Initialize(_ context.Context) error { + if _, err := exec.LookPath("docker"); err != nil { + return fmt.Errorf("docker CLI not found in PATH") + } + slog.Info("docker CLI found") + return nil +} + +func (d *DockerTool) Shutdown(_ context.Context) error { return nil } + +func (d *DockerTool) HealthCheck(ctx context.Context) error { + return d.dockerCmd(ctx, "version").Run() +} + +func (d *DockerTool) Register(mcpServer *server.MCPServer) error { + mcpServer.AddTool(mcp.NewTool("docker_ps", + mcp.WithDescription("列出 Docker 容器"), + mcp.WithString("filter", mcp.Description("按名称过滤(支持正则)")), + mcp.WithBoolean("all", mcp.Description("是否包含已停止的容器,默认 false")), + ), d.handleList) + + mcpServer.AddTool(mcp.NewTool("docker_logs", + mcp.WithDescription("获取容器日志(受白名单和大小限制保护)"), + mcp.WithString("container", mcp.Required(), mcp.Description("容器名称或 ID")), + mcp.WithNumber("tail", mcp.Description("返回最后 N 行,默认 100")), + mcp.WithString("since", mcp.Description("从多久前开始,如 15m、1h,默认 15m")), + ), d.handleLogs) + + mcpServer.AddTool(mcp.NewTool("docker_inspect", + mcp.WithDescription("查看容器详细信息(受白名单保护)"), + mcp.WithString("container", mcp.Required(), mcp.Description("容器名称或 ID")), + ), d.handleInspect) + + return nil +} + +// --- 安全校验 --- + +func (d *DockerTool) validateContainer(name string) error { + if !d.containerRegex.MatchString(name) { + return fmt.Errorf("container %q not in allowed pattern", name) + } + return nil +} + +// --- docker CLI 封装 --- + +func (d *DockerTool) dockerCmd(ctx context.Context, args ...string) *exec.Cmd { + return exec.CommandContext(ctx, "docker", args...) +} + +func (d *DockerTool) dockerOutput(ctx context.Context, args ...string) ([]byte, error) { + cmd := d.dockerCmd(ctx, args...) + var stderr bytes.Buffer + cmd.Stderr = &stderr + + output, err := cmd.Output() + if err != nil { + return nil, fmt.Errorf("docker %s: %v\n%s", strings.Join(args, " "), err, stderr.String()) + } + return output, nil +} + +// --- Handlers --- + +func (d *DockerTool) handleList(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { + args := getArgs(req) + filterName, _ := args["filter"].(string) + showAll, _ := args["all"].(bool) + + dockerArgs := []string{"ps", "--format", "{{json .}}"} + if showAll { + dockerArgs = append(dockerArgs, "-a") + } + + output, err := d.dockerOutput(ctx, dockerArgs...) + if err != nil { + return mcp.NewToolResultError(err.Error()), nil + } + + type dockerPsRow struct { + Names string `json:"Names"` + Image string `json:"Image"` + Status string `json:"Status"` + CreatedAt string `json:"CreatedAt"` + Ports string `json:"Ports"` + } + + var result []dockerPsRow + for _, line := range strings.Split(strings.TrimSpace(string(output)), "\n") { + if line == "" { + continue + } + var row dockerPsRow + if err := json.Unmarshal([]byte(line), &row); err != nil { + continue + } + + // 应用名称过滤 + if filterName != "" { + re, err := regexp.Compile(filterName) + if err != nil { + continue + } + if !re.MatchString(row.Names) { + continue + } + } + + // 白名单检查 + if err := d.validateContainer(row.Names); err != nil { + continue + } + + result = append(result, row) + } + + data, _ := json.MarshalIndent(result, "", " ") + return mcp.NewToolResultText(fmt.Sprintf("```json\n%s\n```", string(data))), nil +} + +func (d *DockerTool) handleLogs(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { + args := getArgs(req) + containerName, _ := args["container"].(string) + + if err := d.validateContainer(containerName); err != nil { + return mcp.NewToolResultError(err.Error()), nil + } + + tail := "100" + if t, ok := args["tail"].(float64); ok { + tail = fmt.Sprintf("%d", int(t)) + } + + since := "15m" + if s, ok := args["since"].(string); ok && s != "" { + dur, err := time.ParseDuration(s) + if err == nil && d.maxLogSince > 0 && dur > d.maxLogSince { + since = d.maxLogSince.String() + } else { + since = s + } + } + + output, err := d.dockerOutput(ctx, "logs", "--tail", tail, "--since", since, "--timestamps", containerName) + if err != nil { + // docker logs 返回非 0 时容器可能不存在 + return mcp.NewToolResultError(fmt.Sprintf("docker logs failed: %v\nOutput: %s", err, string(output))), nil + } + + // 截断字节数 + text := string(output) + if len(text) > int(d.maxLogBytes) { + text = text[len(text)-int(d.maxLogBytes):] + } + + // 截断行数 + lines := strings.Split(text, "\n") + if len(lines) > d.maxLogLines { + lines = lines[len(lines)-d.maxLogLines:] + } + + return mcp.NewToolResultText(strings.Join(lines, "\n")), nil +} + +func (d *DockerTool) handleInspect(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { + args := getArgs(req) + containerName, _ := args["container"].(string) + + if err := d.validateContainer(containerName); err != nil { + return mcp.NewToolResultError(err.Error()), nil + } + + output, err := d.dockerOutput(ctx, "inspect", containerName) + if err != nil { + return mcp.NewToolResultError(err.Error()), nil + } + + // 只提取安全字段 + var inspects []struct { + Name string `json:"Name"` + ID string `json:"Id"` + Image string `json:"Image"` + State struct { + Status string `json:"Status"` + Running bool `json:"Running"` + StartedAt string `json:"StartedAt"` + Pid int `json:"Pid"` + } `json:"State"` + Created string `json:"Created"` + Config struct { + Image string `json:"Image"` + Env []string `json:"Env"` + } `json:"Config"` + Mounts []struct { + Source string `json:"Source"` + Destination string `json:"Destination"` + Mode string `json:"Mode"` + } `json:"Mounts"` + } + + if err := json.Unmarshal(output, &inspects); err != nil { + return mcp.NewToolResultError(fmt.Sprintf("parse inspect: %v", err)), nil + } + + if len(inspects) == 0 { + return mcp.NewToolResultError("container not found"), nil + } + + insp := inspects[0] + // 不暴露 Env(含密钥),只暴露安全信息 + safe := map[string]any{ + "name": strings.TrimPrefix(insp.Name, "/"), + "id": insp.ID[:12], + "image": insp.Config.Image, + "state": insp.State, + "created": insp.Created, + "mounts": insp.Mounts, + } + + data, _ := json.MarshalIndent(safe, "", " ") + return mcp.NewToolResultText(fmt.Sprintf("```json\n%s\n```", string(data))), nil +} diff --git a/internal/tool/system.go b/internal/tool/system.go new file mode 100644 index 0000000..4869382 --- /dev/null +++ b/internal/tool/system.go @@ -0,0 +1,82 @@ +package tool + +import ( + "context" + "encoding/json" + "fmt" + "runtime" + "time" + + "github.com/mark3labs/mcp-go/mcp" + "github.com/mark3labs/mcp-go/server" +) + +var StartTime = time.Now() + +// SystemTool 提供 health 和 info 两个基础工具。 +type SystemTool struct { + version string + registry HealthChecker +} + +// HealthChecker 用于 system 工具检查其他工具的连通性。 +type HealthChecker interface { + HealthCheckAll(ctx context.Context) map[string]error +} + +func NewSystemTool(version string, reg HealthChecker) *SystemTool { + return &SystemTool{version: version, registry: reg} +} + +func (s *SystemTool) Name() string { return "system" } +func (s *SystemTool) Description() string { return "系统健康检查和信息查询" } + +func (s *SystemTool) Initialize(_ context.Context) error { return nil } +func (s *SystemTool) Shutdown(_ context.Context) error { return nil } +func (s *SystemTool) HealthCheck(_ context.Context) error { return nil } + +func (s *SystemTool) Register(mcpServer *server.MCPServer) error { + mcpServer.AddTool(mcp.NewTool("health", + mcp.WithDescription("服务健康检查:运行时间、内存使用、连接状态"), + ), s.handleHealth) + + mcpServer.AddTool(mcp.NewTool("info", + mcp.WithDescription("服务信息:版本号、已配置的数据库、已加载的工具"), + ), s.handleInfo) + + return nil +} + +func (s *SystemTool) handleHealth(ctx context.Context, _ mcp.CallToolRequest) (*mcp.CallToolResult, error) { + var m runtime.MemStats + runtime.ReadMemStats(&m) + + result := map[string]any{ + "status": "ok", + "uptime": time.Since(StartTime).String(), + "goroutines": runtime.NumGoroutine(), + "heap_mb": float64(m.HeapAlloc) / 1024 / 1024, + "gc_cycles": m.NumGC, + } + + if s.registry != nil { + result["connections"] = s.registry.HealthCheckAll(ctx) + } + + data, _ := json.MarshalIndent(result, "", " ") + return mcp.NewToolResultText(string(data)), nil +} + +func (s *SystemTool) handleInfo(ctx context.Context, _ mcp.CallToolRequest) (*mcp.CallToolResult, error) { + result := map[string]any{ + "version": s.version, + "go_version": runtime.Version(), + "os": runtime.GOOS, + "arch": runtime.GOARCH, + "cpus": runtime.NumCPU(), + "start_time": StartTime.Format(time.RFC3339), + } + + data, _ := json.MarshalIndent(result, "", " ") + return mcp.NewToolResultText(fmt.Sprintf("```json\n%s\n```", string(data))), nil +} diff --git a/internal/tool/tool.go b/internal/tool/tool.go new file mode 100644 index 0000000..3bfca9c --- /dev/null +++ b/internal/tool/tool.go @@ -0,0 +1,37 @@ +package tool + +import ( + "context" + + "github.com/mark3labs/mcp-go/mcp" + "github.com/mark3labs/mcp-go/server" +) + +// getArgs 把 req.Params.Arguments 转换为 map[string]any。 +func getArgs(req mcp.CallToolRequest) map[string]any { + if args, ok := req.Params.Arguments.(map[string]any); ok { + return args + } + return map[string]any{} +} + +// Tool 是所有 MCP 工具必须实现的接口。 +// 每个工具模块实现此接口,然后注册到 Registry。 +type Tool interface { + Name() string + + // Description 返回工具的一句话描述 + Description() string + + // Register 向 MCP Server 注册自己的工具处理器 + Register(mcpServer *server.MCPServer) error + + // Initialize 建立连接、初始化资源 + Initialize(ctx context.Context) error + + // Shutdown 释放资源、关闭连接 + Shutdown(ctx context.Context) error + + // HealthCheck 返回当前工具的连接状态,nil 表示健康 + HealthCheck(ctx context.Context) error +} diff --git a/test/integration/mcp_test.go b/test/integration/mcp_test.go new file mode 100644 index 0000000..51f584d --- /dev/null +++ b/test/integration/mcp_test.go @@ -0,0 +1,26 @@ +package integration + +import ( + "testing" +) + +func TestToolInterface(t *testing.T) { + t.Run("system_tool_registration", func(t *testing.T) { + // system 工具是编译时注册的,这里验证接口符合预期 + t.Log("system tool interface verified at compile time") + }) +} + +func TestDatabaseTool(t *testing.T) { + if testing.Short() { + t.Skip("skipping integration test in short mode") + } + t.Log("database integration tests require running PostgreSQL/MySQL") +} + +func TestDockerTool(t *testing.T) { + if testing.Short() { + t.Skip("skipping integration test in short mode") + } + t.Log("docker integration tests require Docker daemon") +}