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>
This commit is contained in:
@@ -0,0 +1,168 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user