fb2c3a6cae
- Tool 接口 + Registry 注册中心,工具模块化插拔 - system 工具:health / info - database 工具:db_query / db_tables / db_table_info / db_explain(四层 SQL 安全) - docker 工具:docker_ps / docker_logs / docker_inspect(白名单 + 大小限制) - middleware:auth(JWT+APIKey)/ ratelimit(令牌桶)/ audit(slog 结构化) - stdio / SSE 双传输模式,Viper 多环境配置 - Dockerfile 多阶段构建(golang:alpine → scratch),~15MB 镜像 - docker-compose.yml 一键部署 + 安全加固(read_only / no-new-privileges / cap_drop) - go build ./... / go vet ./... / go test ./... 全部通过 Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
400 lines
11 KiB
Go
400 lines
11 KiB
Go
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
|
|
}
|
|
}
|