From 093511bbdd683e841d3f015b4aad531437b676af Mon Sep 17 00:00:00 2001 From: yangzhaohan Date: Mon, 6 Jul 2026 10:42:23 +0800 Subject: [PATCH] =?UTF-8?q?=E7=B2=BE=E7=AE=80=E9=A1=B9=E7=9B=AE=EF=BC=9A?= =?UTF-8?q?=E5=8F=AA=E4=BF=9D=E7=95=99=20docker=5Flogs=20+=20health=20+=20?= =?UTF-8?q?info=20=E4=B8=89=E4=B8=AA=E5=B7=A5=E5=85=B7?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 移除 database/middleware 等复杂模块,先聚焦核心功能 - docker_logs:封装 docker logs --tail N - 修复 go.mod 版本为 1.24 匹配 Docker 构建镜像 - Dockerfile 改用 alpine + docker-cli,避免 scratch 无 shell 问题 - docker-compose.yml 简化为单服务定义 Co-Authored-By: Claude Opus 4.7 (1M context) --- .env.example | 5 +- Dockerfile | 9 +- Makefile | 37 +-- README.md | 71 ++---- cmd/server/main.go | 70 +----- config/config.go | 67 +----- config/dev.yaml | 24 +- config/prod.yaml | 25 +- config/staging.yaml | 24 -- docker-compose.yml | 47 +--- go.mod | 18 +- go.sum | 45 +--- internal/middleware/audit.go | 42 ---- internal/middleware/auth.go | 94 -------- internal/middleware/ratelimit.go | 53 ---- internal/server/config.go | 17 -- internal/server/registry.go | 17 +- internal/server/server.go | 63 +---- internal/tool/database.go | 399 ------------------------------- internal/tool/docker.go | 272 ++++----------------- internal/tool/system.go | 40 ++-- internal/tool/tool.go | 31 +-- test/integration/mcp_test.go | 26 -- 23 files changed, 144 insertions(+), 1352 deletions(-) delete mode 100644 config/staging.yaml delete mode 100644 internal/middleware/audit.go delete mode 100644 internal/middleware/auth.go delete mode 100644 internal/middleware/ratelimit.go delete mode 100644 internal/server/config.go delete mode 100644 internal/tool/database.go delete mode 100644 test/integration/mcp_test.go diff --git a/.env.example b/.env.example index 91ca60d..5174249 100644 --- a/.env.example +++ b/.env.example @@ -1,4 +1 @@ -OPS_MCP_ENV=prod -OPS_MCP_PORT=8080 -DB_PASSWORD=change_me_in_production -JWT_SECRET=change_me_in_production +OPS_MCP_TRANSPORT=sse diff --git a/Dockerfile b/Dockerfile index d9a23f5..6034d86 100644 --- a/Dockerfile +++ b/Dockerfile @@ -6,15 +6,14 @@ 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 +RUN CGO_ENABLED=0 go build -ldflags="-s -w" -o /server ./cmd/server -FROM scratch +FROM alpine:3.21 + +RUN apk add --no-cache docker-cli ca-certificates 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 index 271b566..17657ff 100644 --- a/Makefile +++ b/Makefile @@ -1,52 +1,33 @@ -.PHONY: build run test lint image clean +.PHONY: build run test image up down logs clean -APP_NAME := ops-mcp -IMAGE_NAME := ops-mcp -GO := go -GOFLAGS := -ldflags="-s -w" - -# ---- Build ---- +APP_NAME := ops-mcp +GO := go build: - $(GO) build $(GOFLAGS) -o $(APP_NAME) ./cmd/server + $(GO) build -ldflags="-s -w" -o $(APP_NAME) ./cmd/server run: build - ./$(APP_NAME) -env=dev - -# ---- Test ---- + ./$(APP_NAME) test: - $(GO) test -race -v ./... - -test-integration: - $(GO) test -race -v ./test/... - -# ---- Lint ---- + $(GO) vet ./... + $(GO) test ./... lint: golangci-lint run ./... -# ---- Docker Compose ---- +image: + docker build -t ops-mcp:latest . 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) diff --git a/README.md b/README.md index 966d35d..528ca39 100644 --- a/README.md +++ b/README.md @@ -1,13 +1,11 @@ # ops-mcp -生产环境 OPS MCP Server —— 安全的数据库查询 + Docker 容器日志工具。 - -Go 编译为 ~15MB 单二进制,scratch 镜像部署。 +Docker 容器日志 MCP Server —— `docker logs` 的 MCP 封装。 ## 快速开始 ```bash -make build && ./ops-mcp -env=dev +make build && ./ops-mcp ``` ## Claude Desktop 配置 @@ -16,65 +14,28 @@ make build && ./ops-mcp -env=dev { "mcpServers": { "ops-mcp": { - "command": "/path/to/ops-mcp", - "args": ["-env=dev"] + "command": "/path/to/ops-mcp" } } } ``` -## Docker Compose 部署(推荐) +## 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 up -d --build ``` -**docker-compose.yml 做了什么:** +然后 Claude Desktop 通过 SSE 连接: -| 项目 | 说明 | -|---|---| -| `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 镜像 +```json +{ "mcpServers": { "ops-mcp": { "url": "http://:8080/sse" } } } ``` + +## MCP 工具 + +| 工具 | 说明 | 参数 | +|---|---|---| +| `docker_logs` | 获取容器日志 | `container`(必填), `tail`(可选,默认 100) | +| `health` | 服务健康检查 | - | +| `info` | 版本信息 | - | diff --git a/cmd/server/main.go b/cmd/server/main.go index 93af605..17ce24f 100644 --- a/cmd/server/main.go +++ b/cmd/server/main.go @@ -1,81 +1,27 @@ package main import ( - "flag" - "fmt" "log/slog" "os" + "ops-mcp/config" "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 { + cfg, err := config.Load() + if err != nil { + slog.Warn("config load failed, using defaults", "err", err) + cfg = &config.Config{Transport: "stdio"} + } + + if err := server.Run(cfg.Transport, tool.NewDockerTool()); 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 index 97375d7..d00ff45 100644 --- a/config/config.go +++ b/config/config.go @@ -1,81 +1,34 @@ package config import ( - "fmt" - "time" + "log/slog" "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) { +func Load() (*Config, error) { v := viper.New() - - v.SetConfigName(env) // dev.yaml / staging.yaml / prod.yaml + v.SetConfigName("config") v.SetConfigType("yaml") v.AddConfigPath("./config") v.AddConfigPath(".") v.AddConfigPath("/config") - v.SetEnvPrefix("OPS_MCP") v.AutomaticEnv() + v.SetDefault("transport", "stdio") + // 配置文件不存在不算错误,直接用默认值 + 环境变量 if err := v.ReadInConfig(); err != nil { - return nil, fmt.Errorf("read config: %w", err) + slog.Warn("no config file found, using defaults and env vars", "err", 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) + cfg := &Config{} + if err := v.Unmarshal(cfg); err != nil { + return nil, err } - - return &cfg, nil + return cfg, nil } diff --git a/config/dev.yaml b/config/dev.yaml index 4991589..5b8067f 100644 --- a/config/dev.yaml +++ b/config/dev.yaml @@ -1,23 +1 @@ -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 +transport: stdio diff --git a/config/prod.yaml b/config/prod.yaml index 6312419..d5ed429 100644 --- a/config/prod.yaml +++ b/config/prod.yaml @@ -1,24 +1 @@ -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 +transport: sse diff --git a/config/staging.yaml b/config/staging.yaml deleted file mode 100644 index 4402173..0000000 --- a/config/staging.yaml +++ /dev/null @@ -1,24 +0,0 @@ -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 index e0f3025..5da0fc4 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -1,5 +1,3 @@ -version: "3.8" - services: ops-mcp: build: @@ -8,50 +6,7 @@ services: 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} + - OPS_MCP_TRANSPORT=sse 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 index e63bcae..4704163 100644 --- a/go.mod +++ b/go.mod @@ -1,31 +1,21 @@ module ops-mcp -go 1.25.0 +go 1.24 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 @@ -35,9 +25,7 @@ require ( 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 + golang.org/x/sys v0.29.0 // indirect + golang.org/x/text v0.28.0 // indirect gopkg.in/yaml.v3 v3.0.1 // indirect ) diff --git a/go.sum b/go.sum index 747d247..67b2913 100644 --- a/go.sum +++ b/go.sum @@ -1,36 +1,21 @@ -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/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI= +github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY= 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= @@ -44,8 +29,8 @@ github.com/pelletier/go-toml/v2 v2.2.4 h1:mye9XuhQ6gvn5h28+VilKrrPoQVanw5PMw/TB0 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/rogpeppe/go-internal v1.9.0 h1:73kH8U+JUqXU8lRuOHeVHaa/SZPifC7BkcraZVejAe8= +github.com/rogpeppe/go-internal v1.9.0/go.mod h1:WtVeX8xhTBvf0smdhujwtBcq4Qrzq/fJaraNFVN+nFs= 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= @@ -58,9 +43,6 @@ 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= @@ -71,19 +53,12 @@ github.com/yosida95/uritemplate/v3 v3.0.2 h1:Ed3Oyj9yrmi9087+NczuL5BwkIc4wvTb5zI 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= +golang.org/x/sys v0.29.0 h1:TPYlXGxvx1MGTn2GiZDhnjPA9wZzZeGKHHmKhHYvgaU= +golang.org/x/sys v0.29.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA= +golang.org/x/text v0.28.0 h1:rhazDwis8INMIwQ4tpjLDzUhx6RlXqZNPEM0huQojng= +golang.org/x/text v0.28.0/go.mod h1:U8nCwOR8jO/marOQ0QbDiOngZVEBB7MAiitBuMjXiNU= 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/check.v1 v1.0.0-20190902080502-41f04d3bba15 h1:YR8cESwS4TdDjEe65xsg0ogRM/Nc3DYOhEAlW+xobZo= +gopkg.in/check.v1 v1.0.0-20190902080502-41f04d3bba15/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= 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 deleted file mode 100644 index 22c171f..0000000 --- a/internal/middleware/audit.go +++ /dev/null @@ -1,42 +0,0 @@ -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 deleted file mode 100644 index f014146..0000000 --- a/internal/middleware/auth.go +++ /dev/null @@ -1,94 +0,0 @@ -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 deleted file mode 100644 index a3155b1..0000000 --- a/internal/middleware/ratelimit.go +++ /dev/null @@ -1,53 +0,0 @@ -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 deleted file mode 100644 index 8ce657f..0000000 --- a/internal/server/config.go +++ /dev/null @@ -1,17 +0,0 @@ -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 index 5a7abcd..9a61e01 100644 --- a/internal/server/registry.go +++ b/internal/server/registry.go @@ -19,18 +19,16 @@ 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()) + slog.Info("tool registered", "name", name) return nil } -// InitializeAll 调用所有工具的 Initialize,任一失败则终止。 func (r *Registry) InitializeAll(ctx context.Context) error { for name, t := range r.tools { slog.Info("initializing tool", "name", name) @@ -41,18 +39,15 @@ func (r *Registry) InitializeAll(ctx context.Context) error { 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 { @@ -64,7 +59,6 @@ func (r *Registry) ShutdownAll(ctx context.Context) []error { 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 { @@ -72,12 +66,3 @@ func (r *Registry) HealthCheckAll(ctx context.Context) map[string]error { } 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 index 616a7d8..5cade1c 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -4,7 +4,6 @@ import ( "context" "fmt" "log/slog" - "net/http" "os" "os/signal" "syscall" @@ -16,12 +15,11 @@ import ( const version = "0.1.0" -// Run 启动 MCP Server,处理信号优雅关闭。 -func Run(cfg *Config, tools ...tool.Tool) error { +// Run 启动 MCP Server。 +func Run(transport string, 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() { @@ -30,7 +28,6 @@ func Run(cfg *Config, tools ...tool.Tool) error { cancel() }() - // 创建注册中心 registry := NewRegistry() for _, t := range tools { if err := registry.Register(t); err != nil { @@ -38,13 +35,11 @@ func Run(cfg *Config, tools ...tool.Tool) error { } } - // 系统工具总是最后注册,确保它能看到所有工具 sysTool := tool.NewSystemTool(version, registry) if err := registry.Register(sysTool); err != nil { - return fmt.Errorf("register system tool: %w", err) + return fmt.Errorf("register system: %w", err) } - // 初始化所有工具 if err := registry.InitializeAll(ctx); err != nil { return fmt.Errorf("initialize: %w", err) } @@ -54,58 +49,20 @@ func Run(cfg *Config, tools ...tool.Tool) error { } }() - // 创建 MCP Server - mcpServer := server.NewMCPServer( - "ops-mcp", - version, - server.WithLogging(), - ) + 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 { + switch transport { case "stdio": - return runStdio(ctx, mcpServer) + slog.Info("starting MCP server", "transport", "stdio") + return server.ServeStdio(mcpServer) case "sse": - return runSSE(ctx, mcpServer, cfg) + slog.Info("starting MCP server", "transport", "sse", "addr", ":8080") + return server.NewSSEServer(mcpServer).Start(":8080") default: - return fmt.Errorf("unknown transport: %s (expect stdio or sse)", cfg.Server.Transport) + return fmt.Errorf("unknown transport: %s", 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 deleted file mode 100644 index 49135ec..0000000 --- a/internal/tool/database.go +++ /dev/null @@ -1,399 +0,0 @@ -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 index 349f77e..7b8d326 100644 --- a/internal/tool/docker.go +++ b/internal/tool/docker.go @@ -3,276 +3,86 @@ package tool import ( "bytes" "context" - "encoding/json" "fmt" "log/slog" "os/exec" - "regexp" - "strings" - "time" + "strconv" "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 -} +// DockerTool 提供 docker logs 查询。 +type DockerTool struct{} -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 NewDockerTool() *DockerTool { + return &DockerTool{} } func (d *DockerTool) Name() string { return "docker" } -func (d *DockerTool) Description() 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") + return fmt.Errorf("docker CLI not found: %w", err) } - slog.Info("docker CLI found") + slog.Info("docker CLI ready") 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() + return exec.CommandContext(ctx, "docker", "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")), + mcp.WithDescription("获取 Docker 容器日志,等价于 docker logs --tail N "), + mcp.WithString("container", + mcp.Required(), + mcp.Description("容器名称或 ID"), + ), + mcp.WithNumber("tail", + mcp.Description("返回最后 N 行日志,默认 100"), + ), ), 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 + container, _ := args["container"].(string) + if container == "" { + return mcp.NewToolResultError("container 参数必填"), nil } - tail := "100" - if t, ok := args["tail"].(float64); ok { - tail = fmt.Sprintf("%d", int(t)) + tail := 100 + if t, ok := args["tail"].(float64); ok && t > 0 { + tail = 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 - } + cmd := exec.CommandContext(ctx, "docker", "logs", + "--tail", strconv.Itoa(tail), + "--timestamps", + container, + ) + + var stdout, stderr bytes.Buffer + cmd.Stdout = &stdout + cmd.Stderr = &stderr + + if err := cmd.Run(); err != nil { + return mcp.NewToolResultError( + fmt.Sprintf("docker logs 失败: %v\n%s", err, stderr.String()), + ), nil } - 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 + output := stdout.String() + if output == "" { + output = "(容器没有日志输出)" } - // 截断字节数 - 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 + return mcp.NewToolResultText(output), nil } diff --git a/internal/tool/system.go b/internal/tool/system.go index 4869382..fa6f263 100644 --- a/internal/tool/system.go +++ b/internal/tool/system.go @@ -13,35 +13,33 @@ import ( 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 } +type SystemTool struct { + version string + registry HealthChecker +} + 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) 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("服务健康检查:运行时间、内存使用、连接状态"), + mcp.WithDescription("服务健康检查:运行时间、内存、连接状态"), ), s.handleHealth) mcpServer.AddTool(mcp.NewTool("info", - mcp.WithDescription("服务信息:版本号、已配置的数据库、已加载的工具"), + mcp.WithDescription("服务信息:版本号、Go 版本、已加载工具"), ), s.handleInfo) return nil @@ -52,11 +50,10 @@ func (s *SystemTool) handleHealth(ctx context.Context, _ mcp.CallToolRequest) (* 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, + "status": "ok", + "uptime": time.Since(StartTime).String(), + "goroutines": runtime.NumGoroutine(), + "heap_mb": float64(m.HeapAlloc) / 1024 / 1024, } if s.registry != nil { @@ -69,14 +66,13 @@ func (s *SystemTool) handleHealth(ctx context.Context, _ mcp.CallToolRequest) (* 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), + "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 index 3bfca9c..73e63d3 100644 --- a/internal/tool/tool.go +++ b/internal/tool/tool.go @@ -7,6 +7,16 @@ import ( "github.com/mark3labs/mcp-go/server" ) +// Tool 是所有 MCP 工具必须实现的接口。 +type Tool interface { + Name() string + Description() string + Register(mcpServer *server.MCPServer) error + Initialize(ctx context.Context) error + Shutdown(ctx context.Context) error + HealthCheck(ctx context.Context) error +} + // 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 { @@ -14,24 +24,3 @@ func getArgs(req mcp.CallToolRequest) map[string]any { } 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 deleted file mode 100644 index 51f584d..0000000 --- a/test/integration/mcp_test.go +++ /dev/null @@ -1,26 +0,0 @@ -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") -}