mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-02 03:22:09 +00:00
Tools: list_tables, describe_table, select_table, insert_into_table, update_table, delete_from_table, list_functions, call_function. Visibility follows the model rules. Filter-based update/delete require filters (never dropped silently), cap the matched rows (MaxWriteRows), support dry_run, and need a single-use confirm token bound to caller, table, filters, data and the matched rows. RegisterFunction adds Go-callback and SQL-procedure functions run in a transaction with BeforeCall/AfterCall hooks. Per-model tools and resources are removed. fix(pgsql): UPDATE with SET and a multi-placeholder WHERE renumbered the WHERE parameters wrongly ($1, $2 became $3, $2); shift them in one pass.
281 lines
7.8 KiB
Go
281 lines
7.8 KiB
Go
package resolvemcp
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"fmt"
|
|
"math"
|
|
"regexp"
|
|
"sort"
|
|
"strings"
|
|
"sync"
|
|
|
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
|
)
|
|
|
|
// Parameter types a function may declare.
|
|
const (
|
|
ParamString = "string"
|
|
ParamInteger = "integer"
|
|
ParamNumber = "number"
|
|
ParamBoolean = "boolean"
|
|
ParamObject = "object"
|
|
ParamArray = "array"
|
|
)
|
|
|
|
// FunctionParam declares one argument of a registered function.
|
|
type FunctionParam struct {
|
|
Name string
|
|
Type string // one of the Param* constants
|
|
Description string
|
|
Required bool
|
|
}
|
|
|
|
// FunctionFunc is a Go function callable through call_function. tx is the transaction the call
|
|
// runs in (OnTxBegin already fired on it); args are validated against the declared params.
|
|
type FunctionFunc func(ctx context.Context, tx common.Database, args map[string]any) (any, error)
|
|
|
|
// Function is a callable registered with Handler.RegisterFunction. Exactly one of Handler (a Go
|
|
// callback) and Procedure (a SQL function called by name) is set.
|
|
type Function struct {
|
|
Name string
|
|
Description string
|
|
Params []FunctionParam
|
|
|
|
// Handler is the Go callback.
|
|
Handler FunctionFunc
|
|
// Procedure is the SQL function to call, optionally schema-qualified. Declared params are
|
|
// passed positionally in declaration order; an omitted optional param is NULL. Rows
|
|
// returned are the result.
|
|
Procedure string
|
|
|
|
// Authorize, when set, decides per caller whether the function is listed and callable.
|
|
// Return nil to allow.
|
|
Authorize func(ctx context.Context) error
|
|
}
|
|
|
|
var (
|
|
functionNameRe = regexp.MustCompile(`^[A-Za-z][A-Za-z0-9_]{0,63}$`)
|
|
procedureRe = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*(\.[A-Za-z_][A-Za-z0-9_]*)?$`)
|
|
paramNameRe = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]{0,63}$`)
|
|
)
|
|
|
|
type functionRegistry struct {
|
|
mu sync.RWMutex
|
|
funcs map[string]Function
|
|
}
|
|
|
|
// RegisterFunction makes f callable through the call_function meta tool. Only registered
|
|
// functions are callable; there is no way to reach an arbitrary SQL function.
|
|
func (h *Handler) RegisterFunction(f Function) error {
|
|
if !functionNameRe.MatchString(f.Name) {
|
|
return fmt.Errorf("resolvemcp: invalid function name %q", f.Name)
|
|
}
|
|
if (f.Handler == nil) == (f.Procedure == "") {
|
|
return fmt.Errorf("resolvemcp: function %q needs exactly one of Handler and Procedure", f.Name)
|
|
}
|
|
if f.Procedure != "" && !procedureRe.MatchString(f.Procedure) {
|
|
return fmt.Errorf("resolvemcp: invalid procedure name %q", f.Procedure)
|
|
}
|
|
seen := map[string]bool{}
|
|
for _, p := range f.Params {
|
|
if !paramNameRe.MatchString(p.Name) || seen[p.Name] {
|
|
return fmt.Errorf("resolvemcp: function %q: invalid or duplicate param %q", f.Name, p.Name)
|
|
}
|
|
seen[p.Name] = true
|
|
switch p.Type {
|
|
case ParamString, ParamInteger, ParamNumber, ParamBoolean, ParamObject, ParamArray:
|
|
default:
|
|
return fmt.Errorf("resolvemcp: function %q param %q: unknown type %q", f.Name, p.Name, p.Type)
|
|
}
|
|
}
|
|
h.functions.mu.Lock()
|
|
defer h.functions.mu.Unlock()
|
|
if h.functions.funcs == nil {
|
|
h.functions.funcs = map[string]Function{}
|
|
}
|
|
if _, dup := h.functions.funcs[f.Name]; dup {
|
|
return fmt.Errorf("resolvemcp: function %q already registered", f.Name)
|
|
}
|
|
h.functions.funcs[f.Name] = f
|
|
return nil
|
|
}
|
|
|
|
func (h *Handler) function(name string) (Function, bool) {
|
|
h.functions.mu.RLock()
|
|
defer h.functions.mu.RUnlock()
|
|
f, ok := h.functions.funcs[name]
|
|
return f, ok
|
|
}
|
|
|
|
// visibleFunctions returns the functions the caller may call, sorted by name.
|
|
func (h *Handler) visibleFunctions(ctx context.Context) []Function {
|
|
h.functions.mu.RLock()
|
|
out := make([]Function, 0, len(h.functions.funcs))
|
|
for _, f := range h.functions.funcs {
|
|
out = append(out, f)
|
|
}
|
|
h.functions.mu.RUnlock()
|
|
sort.Slice(out, func(i, j int) bool { return out[i].Name < out[j].Name })
|
|
visible := out[:0]
|
|
for _, f := range out {
|
|
if f.Authorize == nil || f.Authorize(ctx) == nil {
|
|
visible = append(visible, f)
|
|
}
|
|
}
|
|
return visible
|
|
}
|
|
|
|
// validateArgs checks args against the declared params: required present, no undeclared names,
|
|
// types match. The returned map is what the function receives.
|
|
func validateArgs(params []FunctionParam, args map[string]any) (map[string]any, error) {
|
|
decl := make(map[string]FunctionParam, len(params))
|
|
for _, p := range params {
|
|
decl[p.Name] = p
|
|
}
|
|
for name := range args {
|
|
if _, ok := decl[name]; !ok {
|
|
return nil, invalidArg("unknown argument %q", truncate(name))
|
|
}
|
|
}
|
|
out := make(map[string]any, len(args))
|
|
for _, p := range params {
|
|
v, ok := args[p.Name]
|
|
if !ok || v == nil {
|
|
if p.Required {
|
|
return nil, invalidArg("missing required argument %q", p.Name)
|
|
}
|
|
continue
|
|
}
|
|
if !argMatches(p.Type, v) {
|
|
return nil, invalidArg("argument %q must be %s", p.Name, p.Type)
|
|
}
|
|
out[p.Name] = v
|
|
}
|
|
return out, nil
|
|
}
|
|
|
|
func truncate(s string) string {
|
|
if len(s) > maxKeyEcho {
|
|
return s[:maxKeyEcho] + "..."
|
|
}
|
|
return s
|
|
}
|
|
|
|
func argMatches(typ string, v any) bool {
|
|
switch typ {
|
|
case ParamString:
|
|
_, ok := v.(string)
|
|
return ok
|
|
case ParamBoolean:
|
|
_, ok := v.(bool)
|
|
return ok
|
|
case ParamNumber:
|
|
return isNumber(v)
|
|
case ParamInteger:
|
|
switch n := v.(type) {
|
|
case float64:
|
|
return n == math.Trunc(n) && !math.IsInf(n, 0)
|
|
case int, int64:
|
|
return true
|
|
}
|
|
return false
|
|
case ParamObject:
|
|
_, ok := v.(map[string]any)
|
|
return ok
|
|
case ParamArray:
|
|
_, ok := v.([]any)
|
|
return ok
|
|
}
|
|
return false
|
|
}
|
|
|
|
func isNumber(v any) bool {
|
|
switch v.(type) {
|
|
case float64, int, int64:
|
|
return true
|
|
}
|
|
return false
|
|
}
|
|
|
|
// callProcedure runs f.Procedure on tx.
|
|
func callProcedure(ctx context.Context, tx common.Database, f Function, args map[string]any) (any, error) {
|
|
placeholders := make([]string, len(f.Params))
|
|
values := make([]any, len(f.Params))
|
|
for i, p := range f.Params {
|
|
ph := fmt.Sprintf("$%d", i+1)
|
|
if p.Type == ParamObject || p.Type == ParamArray {
|
|
ph += "::jsonb"
|
|
}
|
|
if v, ok := args[p.Name]; !ok {
|
|
values[i] = nil
|
|
} else if p.Type == ParamObject || p.Type == ParamArray {
|
|
b, err := json.Marshal(v)
|
|
if err != nil {
|
|
return nil, invalidArg("argument %q cannot be encoded", p.Name)
|
|
}
|
|
values[i] = string(b)
|
|
} else {
|
|
values[i] = v
|
|
}
|
|
placeholders[i] = ph
|
|
}
|
|
var rows []map[string]any
|
|
query := fmt.Sprintf("SELECT * FROM %s(%s)", f.Procedure, strings.Join(placeholders, ", "))
|
|
if err := tx.Query(ctx, &rows, query, values...); err != nil {
|
|
return nil, err
|
|
}
|
|
return rows, nil
|
|
}
|
|
|
|
// executeCall validates and runs a registered function in a transaction: BeforeHandle (auth and
|
|
// rules), Authorize, then OnTxBegin, BeforeCall, the function, AfterCall.
|
|
func (h *Handler) executeCall(ctx context.Context, name string, rawArgs map[string]any) (_ any, retErr error) {
|
|
defer recoverPanic(&retErr)
|
|
ctx, cancel := h.callContext(ctx)
|
|
defer cancel()
|
|
|
|
f, ok := h.function(name)
|
|
if !ok {
|
|
return nil, invalidArg("unknown function %q", truncate(name))
|
|
}
|
|
hookCtx := &HookContext{Context: ctx, Handler: h, Entity: name, Operation: "call_function", Tx: h.db}
|
|
if err := h.hooks.Execute(BeforeHandle, hookCtx); err != nil {
|
|
return nil, err
|
|
}
|
|
// Same answer for "not authorized" and "does not exist" so names cannot be probed.
|
|
if f.Authorize != nil && f.Authorize(ctx) != nil {
|
|
return nil, invalidArg("unknown function %q", truncate(name))
|
|
}
|
|
args, err := validateArgs(f.Params, rawArgs)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
hookCtx.Data = args
|
|
|
|
var result any
|
|
err = h.runInTx(ctx, hookCtx, func(tx common.Database) error {
|
|
if err := h.hooks.Execute(BeforeCall, hookCtx); err != nil {
|
|
return err
|
|
}
|
|
if m, ok := hookCtx.Data.(map[string]any); ok {
|
|
args = m
|
|
}
|
|
var err error
|
|
if f.Handler != nil {
|
|
result, err = f.Handler(ctx, tx, args)
|
|
} else {
|
|
result, err = callProcedure(ctx, tx, f, args)
|
|
}
|
|
if err != nil {
|
|
return err
|
|
}
|
|
hookCtx.Result = result
|
|
return h.hooks.Execute(AfterCall, hookCtx)
|
|
})
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return result, nil
|
|
}
|