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:
@@ -90,7 +90,7 @@ func main() {
|
||||
go b.RequestConfigWithRetry(ctx, contract.ConfigKindEmbedding, applyEmbed)
|
||||
go b.RequestConfigWithRetry(ctx, contract.ConfigKindChat, applyChat)
|
||||
|
||||
gw := mcp.NewGateway(b, engine, mem, hist, ragEngine)
|
||||
gw := mcp.NewGateway(b, engine, mem, hist, ragEngine, pgDSN)
|
||||
|
||||
log.Println("[mcp_go] serving MCP over sundynix.tools.go.* (Ctrl-C to quit)")
|
||||
if err := gw.Serve(ctx); err != nil && err != context.Canceled {
|
||||
|
||||
@@ -3,6 +3,7 @@ package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log"
|
||||
@@ -10,6 +11,7 @@ import (
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
sharedbus "github.com/sundynix/sundynix-shared/bus"
|
||||
@@ -30,6 +32,11 @@ type Gateway struct {
|
||||
history *history.Store
|
||||
rag *rag.Engine
|
||||
tools map[string]toolDef // 工具注册表:唯一事实源,dispatch 与 list_tools 共用,杜绝漂移
|
||||
|
||||
pgDSN string // 平台 PG DSN(sql_query 兜底库;SQL_QUERY_DSN 未设时用它)
|
||||
sqlOnce sync.Once // sql_query 只读连接懒连一次
|
||||
sqlDB *sql.DB //
|
||||
sqlDBErr error //
|
||||
}
|
||||
|
||||
// paramSpec 是一个工具参数的声明(供自主 agent 据此生成调用入参)。
|
||||
@@ -53,8 +60,8 @@ type toolDef struct {
|
||||
handler func(context.Context, *contract.ToolCall) *contract.ToolResult
|
||||
}
|
||||
|
||||
func NewGateway(b *sharedbus.Bus, s *search.Hybrid, m *memory.Store, h *history.Store, r *rag.Engine) *Gateway {
|
||||
g := &Gateway{bus: b, search: s, memory: m, history: h, rag: r}
|
||||
func NewGateway(b *sharedbus.Bus, s *search.Hybrid, m *memory.Store, h *history.Store, r *rag.Engine, pgDSN string) *Gateway {
|
||||
g := &Gateway{bus: b, search: s, memory: m, history: h, rag: r, pgDSN: pgDSN}
|
||||
g.tools = g.buildRegistry()
|
||||
return g
|
||||
}
|
||||
@@ -129,6 +136,12 @@ func (g *Gateway) buildRegistry() map[string]toolDef {
|
||||
params: []paramSpec{{Name: "tz", Type: "string", Desc: "可选时区,如 Asia/Shanghai;缺省服务器本地时区"}},
|
||||
handler: g.currentDatetime,
|
||||
},
|
||||
"sql_query": {
|
||||
cn: "SQL查询", desc: "对数据库执行只读 SQL 查询(仅 SELECT/WITH),返回结果表。需要查业务数据/统计时调用。",
|
||||
agent: true,
|
||||
params: []paramSpec{{Name: "sql", Type: "string", Desc: "只读 SQL,如 SELECT count(*) FROM sundynix_task", Required: true}},
|
||||
handler: g.sqlQuery,
|
||||
},
|
||||
|
||||
// —— 仅内部/流水线/管理用,不暴露给自主 agent ——
|
||||
"kb_ingest": {cn: "知识入库", desc: "文本切块 → 向量化 → 写入 Milvus / Bleve", handler: g.kbIngest},
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -56,3 +56,32 @@ func TestDDGUnwrap(t *testing.T) {
|
||||
t.Fatalf("ddgUnwrap 解码错误: %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateReadOnlySQL(t *testing.T) {
|
||||
okCases := []string{
|
||||
"SELECT 1",
|
||||
"select count(*) from sundynix_task",
|
||||
"WITH x AS (SELECT 1) SELECT * FROM x",
|
||||
" select * from t where created_at > now() ", // created/update 作为列名一部分不应误伤
|
||||
}
|
||||
for _, q := range okCases {
|
||||
if _, ok := validateReadOnlySQL(q); !ok {
|
||||
t.Errorf("应通过: %q", q)
|
||||
}
|
||||
}
|
||||
badCases := []string{
|
||||
"insert into t values (1)",
|
||||
"update t set a=1",
|
||||
"delete from t",
|
||||
"drop table t",
|
||||
"select 1; drop table t",
|
||||
"truncate t",
|
||||
"SELECT 1; SELECT 2",
|
||||
"grant all on t to x",
|
||||
}
|
||||
for _, q := range badCases {
|
||||
if _, ok := validateReadOnlySQL(q); ok {
|
||||
t.Errorf("应拒绝: %q", q)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user