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 } }