ca38dcd0c9
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>
169 lines
5.1 KiB
Go
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)
|
|
}
|
|
}
|