Files
Blizzard ca38dcd0c9 feat(tools): 新增 sql_query 只读 SQL 查询工具(agent 可见)
agent 可对数据库执行只读查询。三重防护:
- 静态校验:仅单条 SELECT/WITH,词边界拒 insert/update/delete/drop/alter/create/
  truncate/grant 等写/DDL 关键字(不误伤 created_at 之类列名)。
- 只读事务(sql.TxOptions{ReadOnly:true}):Postgres 引擎级强制只读,硬兜底。
- 行数(100)+超时(10s)+单元格(200 rune)上限。
连接:SQL_QUERY_DSN 优先(生产应指向专用只读库/账号),未设回退服务已解析的平台 PG DSN
(经 NewGateway 传入 g.pgDSN,不再裸读 env)。独立小连接池(8)。

测试:单测 validateReadOnlySQL(放行 SELECT/WITH/列名含 created;拒写/DDL/多语句)。
live 自主 agent 实测:查 sundynix_task=157 行、sundynix_model=2 行。

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-06-24 17:07:34 +08:00

169 lines
5.1 KiB
Go

package mcp
import (
"context"
"database/sql"
"fmt"
"os"
"strings"
"time"
"github.com/sundynix/sundynix-shared/contract"
"gorm.io/driver/postgres"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
const (
sqlQueryMaxRows = 100 // 结果行上限,防超大结果撑爆上下文
sqlQueryTimeout = 10 * time.Second // 单次查询超时
sqlCellMaxRunes = 200 // 单元格文本上限
)
// sqlQueryDB 懒连只读查询库:优先 SQL_QUERY_DSN(生产应指向专用只读库 / 只读账号),
// 未配置则回退本服务已解析的平台 PG DSN(g.pgDSN,仅供开发;线上务必单配,避免把平台库直接暴露给 agent)。
func (g *Gateway) sqlQueryDB() (*sql.DB, error) {
g.sqlOnce.Do(func() {
dsn := strings.TrimSpace(os.Getenv("SQL_QUERY_DSN"))
if dsn == "" {
dsn = g.pgDSN
}
if dsn == "" {
g.sqlDBErr = fmt.Errorf("未配置 SQL_QUERY_DSN")
return
}
gdb, err := gorm.Open(postgres.New(postgres.Config{DSN: dsn}), &gorm.Config{
Logger: logger.Default.LogMode(logger.Silent),
})
if err != nil {
g.sqlDBErr = err
return
}
db, err := gdb.DB()
if err != nil {
g.sqlDBErr = err
return
}
db.SetMaxOpenConns(8) // 只读查询独立小池,不挤占主连接
db.SetMaxIdleConns(2)
db.SetConnMaxLifetime(time.Hour)
g.sqlDB = db
})
return g.sqlDB, g.sqlDBErr
}
// sqlQuery 对数据库执行只读 SQL。三重防护:① 仅允许单条 SELECT/WITH 语句(拒多语句/写/DDL);
// ② 在「只读事务」中执行(Postgres 引擎级强制只读,是真正的兜底);③ 行数 + 超时 + 单元格上限。
func (g *Gateway) sqlQuery(ctx context.Context, call *contract.ToolCall) *contract.ToolResult {
q := strings.TrimSpace(fmt.Sprint(call.Args["sql"]))
if q == "" || q == "<nil>" {
return &contract.ToolResult{OK: false, Error: "sql_query: sql 必填"}
}
if reason, ok := validateReadOnlySQL(q); !ok {
return &contract.ToolResult{OK: false, Error: "sql_query: " + reason}
}
db, err := g.sqlQueryDB()
if err != nil {
return &contract.ToolResult{OK: false, Error: "sql_query: " + err.Error()}
}
cctx, cancel := context.WithTimeout(ctx, sqlQueryTimeout)
defer cancel()
// 只读事务:即便上面的语句校验被绕过,Postgres 也会在引擎层拒绝任何写操作。
tx, err := db.BeginTx(cctx, &sql.TxOptions{ReadOnly: true})
if err != nil {
return &contract.ToolResult{OK: false, Error: "sql_query: " + err.Error()}
}
defer func() { _ = tx.Rollback() }()
rows, err := tx.QueryContext(cctx, q)
if err != nil {
return &contract.ToolResult{OK: false, Error: "sql_query: " + err.Error()}
}
defer rows.Close()
cols, _ := rows.Columns()
var b strings.Builder
b.WriteString(strings.Join(cols, " | ") + "\n")
b.WriteString(strings.Repeat("-", len(strings.Join(cols, " | "))) + "\n")
n := 0
for rows.Next() {
if n >= sqlQueryMaxRows {
break
}
cells := make([]any, len(cols))
ptrs := make([]any, len(cols))
for i := range cells {
ptrs[i] = &cells[i]
}
if err := rows.Scan(ptrs...); err != nil {
return &contract.ToolResult{OK: false, Error: "sql_query: scan " + err.Error()}
}
strs := make([]string, len(cols))
for i, c := range cells {
strs[i] = truncateRunes(cellToString(c), sqlCellMaxRunes)
}
b.WriteString(strings.Join(strs, " | ") + "\n")
n++
}
if err := rows.Err(); err != nil {
return &contract.ToolResult{OK: false, Error: "sql_query: " + err.Error()}
}
b.WriteString(fmt.Sprintf("(%d 行%s)", n, map[bool]string{true: ",已截断至上限"}[n >= sqlQueryMaxRows]))
return &contract.ToolResult{OK: true, Content: b.String()}
}
// validateReadOnlySQL 静态校验:单条语句、以 SELECT/WITH 开头、不含写/DDL 关键字。
// 与只读事务双保险(这层挡明显误用,事务层是硬保证)。
func validateReadOnlySQL(q string) (reason string, ok bool) {
s := strings.TrimSpace(q)
s = strings.TrimSuffix(s, ";")
if strings.Contains(s, ";") {
return "只允许单条语句", false
}
low := strings.ToLower(s)
if !strings.HasPrefix(low, "select") && !strings.HasPrefix(low, "with") {
return "只允许 SELECT / WITH 查询", false
}
// 词边界匹配写/DDL 关键字(避免误伤列名如 created_at)。
for _, kw := range []string{"insert", "update", "delete", "drop", "alter", "create",
"truncate", "grant", "revoke", "comment", "copy", "merge", "call", "do "} {
if containsWord(low, kw) {
return "检测到写/DDL 关键字「" + strings.TrimSpace(kw) + "」,仅允许只读查询", false
}
}
return "", true
}
func containsWord(s, word string) bool {
for i := 0; ; {
idx := strings.Index(s[i:], word)
if idx < 0 {
return false
}
j := i + idx
before := j == 0 || !isWordChar(rune(s[j-1]))
after := j+len(word) >= len(s) || !isWordChar(rune(s[j+len(word)]))
if before && after {
return true
}
i = j + len(word)
}
}
func isWordChar(r rune) bool {
return r == '_' || (r >= 'a' && r <= 'z') || (r >= 'A' && r <= 'Z') || (r >= '0' && r <= '9')
}
func cellToString(v any) string {
switch x := v.(type) {
case nil:
return "NULL"
case []byte:
return string(x)
case time.Time:
return x.Format("2006-01-02 15:04:05")
default:
return fmt.Sprint(x)
}
}