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.ConfigKindEmbedding, applyEmbed)
|
||||||
go b.RequestConfigWithRetry(ctx, contract.ConfigKindChat, applyChat)
|
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)")
|
log.Println("[mcp_go] serving MCP over sundynix.tools.go.* (Ctrl-C to quit)")
|
||||||
if err := gw.Serve(ctx); err != nil && err != context.Canceled {
|
if err := gw.Serve(ctx); err != nil && err != context.Canceled {
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ package mcp
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"database/sql"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
"log"
|
"log"
|
||||||
@@ -10,6 +11,7 @@ import (
|
|||||||
"path/filepath"
|
"path/filepath"
|
||||||
"sort"
|
"sort"
|
||||||
"strings"
|
"strings"
|
||||||
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
sharedbus "github.com/sundynix/sundynix-shared/bus"
|
sharedbus "github.com/sundynix/sundynix-shared/bus"
|
||||||
@@ -30,6 +32,11 @@ type Gateway struct {
|
|||||||
history *history.Store
|
history *history.Store
|
||||||
rag *rag.Engine
|
rag *rag.Engine
|
||||||
tools map[string]toolDef // 工具注册表:唯一事实源,dispatch 与 list_tools 共用,杜绝漂移
|
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 据此生成调用入参)。
|
// paramSpec 是一个工具参数的声明(供自主 agent 据此生成调用入参)。
|
||||||
@@ -53,8 +60,8 @@ type toolDef struct {
|
|||||||
handler func(context.Context, *contract.ToolCall) *contract.ToolResult
|
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 {
|
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}
|
g := &Gateway{bus: b, search: s, memory: m, history: h, rag: r, pgDSN: pgDSN}
|
||||||
g.tools = g.buildRegistry()
|
g.tools = g.buildRegistry()
|
||||||
return g
|
return g
|
||||||
}
|
}
|
||||||
@@ -129,6 +136,12 @@ func (g *Gateway) buildRegistry() map[string]toolDef {
|
|||||||
params: []paramSpec{{Name: "tz", Type: "string", Desc: "可选时区,如 Asia/Shanghai;缺省服务器本地时区"}},
|
params: []paramSpec{{Name: "tz", Type: "string", Desc: "可选时区,如 Asia/Shanghai;缺省服务器本地时区"}},
|
||||||
handler: g.currentDatetime,
|
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 ——
|
// —— 仅内部/流水线/管理用,不暴露给自主 agent ——
|
||||||
"kb_ingest": {cn: "知识入库", desc: "文本切块 → 向量化 → 写入 Milvus / Bleve", handler: g.kbIngest},
|
"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)
|
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