From ca38dcd0c92a406c240f8aee0c3e7fb1914cd1a2 Mon Sep 17 00:00:00 2001 From: Blizzard Date: Wed, 24 Jun 2026 17:07:34 +0800 Subject: [PATCH] =?UTF-8?q?feat(tools):=20=E6=96=B0=E5=A2=9E=20sql=5Fquery?= =?UTF-8?q?=20=E5=8F=AA=E8=AF=BB=20SQL=20=E6=9F=A5=E8=AF=A2=E5=B7=A5?= =?UTF-8?q?=E5=85=B7=EF=BC=88agent=20=E5=8F=AF=E8=A7=81=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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) --- sundynix-mcp-go/cmd/server/main.go | 2 +- sundynix-mcp-go/internal/mcp/gateway.go | 17 +- sundynix-mcp-go/internal/mcp/sql_query.go | 168 ++++++++++++++++++ .../internal/mcp/tools_extra_test.go | 29 +++ 4 files changed, 213 insertions(+), 3 deletions(-) create mode 100644 sundynix-mcp-go/internal/mcp/sql_query.go diff --git a/sundynix-mcp-go/cmd/server/main.go b/sundynix-mcp-go/cmd/server/main.go index 5bfd161..71b067f 100644 --- a/sundynix-mcp-go/cmd/server/main.go +++ b/sundynix-mcp-go/cmd/server/main.go @@ -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 { diff --git a/sundynix-mcp-go/internal/mcp/gateway.go b/sundynix-mcp-go/internal/mcp/gateway.go index b696a58..eaa259c 100644 --- a/sundynix-mcp-go/internal/mcp/gateway.go +++ b/sundynix-mcp-go/internal/mcp/gateway.go @@ -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}, diff --git a/sundynix-mcp-go/internal/mcp/sql_query.go b/sundynix-mcp-go/internal/mcp/sql_query.go new file mode 100644 index 0000000..4460f36 --- /dev/null +++ b/sundynix-mcp-go/internal/mcp/sql_query.go @@ -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 == "" { + 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) + } +} diff --git a/sundynix-mcp-go/internal/mcp/tools_extra_test.go b/sundynix-mcp-go/internal/mcp/tools_extra_test.go index 96ffd04..d78808a 100644 --- a/sundynix-mcp-go/internal/mcp/tools_extra_test.go +++ b/sundynix-mcp-go/internal/mcp/tools_extra_test.go @@ -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) + } + } +}