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) } }