mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-09 23:06:27 +00:00
feat(resolvemcp): replace per-model tools with fixed meta tools, guarded filter writes and a function registry
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.
This commit is contained in:
@@ -5,7 +5,9 @@ import (
|
|||||||
"database/sql"
|
"database/sql"
|
||||||
"fmt"
|
"fmt"
|
||||||
"reflect"
|
"reflect"
|
||||||
|
"regexp"
|
||||||
"sort"
|
"sort"
|
||||||
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
@@ -767,6 +769,9 @@ func (p *PgSQLInsertQuery) Scan(ctx context.Context, dest interface{}) (err erro
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// placeholderRe matches a numbered SQL parameter such as $12.
|
||||||
|
var placeholderRe = regexp.MustCompile(`\$\d+`)
|
||||||
|
|
||||||
// PgSQLUpdateQuery implements UpdateQuery for PostgreSQL
|
// PgSQLUpdateQuery implements UpdateQuery for PostgreSQL
|
||||||
type PgSQLUpdateQuery struct {
|
type PgSQLUpdateQuery struct {
|
||||||
db *sql.DB
|
db *sql.DB
|
||||||
@@ -897,23 +902,17 @@ func (p *PgSQLUpdateQuery) Exec(ctx context.Context) (res common.Result, err err
|
|||||||
p.tableName,
|
p.tableName,
|
||||||
strings.Join(setClauses, ", "))
|
strings.Join(setClauses, ", "))
|
||||||
|
|
||||||
// Update WHERE clause parameter numbers to continue after SET parameters
|
// WHERE placeholders were numbered from $1 as the clauses were added; shift every one past
|
||||||
|
// the SET parameters in a single pass (replacing one number at a time would rewrite a
|
||||||
|
// number it had just produced, e.g. "$1, $2" -> "$3, $2").
|
||||||
if len(p.whereClauses) > 0 {
|
if len(p.whereClauses) > 0 {
|
||||||
|
shift := len(setArgs)
|
||||||
updatedWhereClauses := make([]string, 0, len(p.whereClauses))
|
updatedWhereClauses := make([]string, 0, len(p.whereClauses))
|
||||||
for _, whereClause := range p.whereClauses {
|
for _, whereClause := range p.whereClauses {
|
||||||
// Find and replace parameter placeholders
|
updatedWhereClauses = append(updatedWhereClauses, placeholderRe.ReplaceAllStringFunc(whereClause, func(m string) string {
|
||||||
updatedClause := whereClause
|
n, _ := strconv.Atoi(m[1:])
|
||||||
paramNum := i
|
return fmt.Sprintf("$%d", n+shift)
|
||||||
// Count how many parameters are in this WHERE clause
|
}))
|
||||||
placeholderCount := strings.Count(whereClause, "$")
|
|
||||||
for j := 0; j < placeholderCount; j++ {
|
|
||||||
oldParam := fmt.Sprintf("$%d", j+1)
|
|
||||||
newParam := fmt.Sprintf("$%d", paramNum)
|
|
||||||
updatedClause = strings.Replace(updatedClause, oldParam, newParam, 1)
|
|
||||||
paramNum++
|
|
||||||
}
|
|
||||||
updatedWhereClauses = append(updatedWhereClauses, updatedClause)
|
|
||||||
i = paramNum
|
|
||||||
}
|
}
|
||||||
p.whereClauses = updatedWhereClauses
|
p.whereClauses = updatedWhereClauses
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -627,3 +627,23 @@ func TestRawSQL(t *testing.T) {
|
|||||||
|
|
||||||
assert.NoError(t, mock.ExpectationsWereMet())
|
assert.NoError(t, mock.ExpectationsWereMet())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// WHERE placeholders must be shifted past the SET parameters without rewriting numbers the
|
||||||
|
// shift itself produced: "a = ? AND b = ?" must stay in order.
|
||||||
|
func TestPgSQLUpdateQuery_WherePlaceholdersAfterSet(t *testing.T) {
|
||||||
|
db, mock, err := sqlmock.New()
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer db.Close()
|
||||||
|
|
||||||
|
mock.ExpectExec(`UPDATE users SET name = \$1 WHERE a = \$2 AND b = \$3 AND "id" IN \(\$4, \$5\)`).
|
||||||
|
WithArgs("n", 10, 20, 7, 8).
|
||||||
|
WillReturnResult(sqlmock.NewResult(0, 2))
|
||||||
|
|
||||||
|
adapter := NewPgSQLAdapter(db)
|
||||||
|
_, err = adapter.NewUpdate().Table("users").SetMap(map[string]interface{}{"name": "n"}).
|
||||||
|
Where("a = ? AND b = ?", 10, 20).
|
||||||
|
Where(`"id" IN (?, ?)`, 7, 8).
|
||||||
|
Exec(context.Background())
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, mock.ExpectationsWereMet())
|
||||||
|
}
|
||||||
|
|||||||
@@ -27,12 +27,13 @@ var allowedPoolHookTx = map[string]int{
|
|||||||
"resolvespec/handler.go": 1,
|
"resolvespec/handler.go": 1,
|
||||||
"websocketspec/handler.go": 1,
|
"websocketspec/handler.go": 1,
|
||||||
"resolvemcp/handler.go": 4,
|
"resolvemcp/handler.go": 4,
|
||||||
|
"resolvemcp/annotation.go": 1, // BeforeHandle context of the annotate tool
|
||||||
|
"resolvemcp/writewhere.go": 1, // BeforeHandle context of filter writes
|
||||||
|
"resolvemcp/functions.go": 1, // BeforeHandle context of call_function
|
||||||
}
|
}
|
||||||
|
|
||||||
// allowedPoolQuery: statements outside the request path.
|
// allowedPoolQuery: statements outside the request path.
|
||||||
var allowedPoolQuery = map[string]int{
|
var allowedPoolQuery = map[string]int{}
|
||||||
"resolvemcp/annotation.go": 2, // tool annotations, not a data request
|
|
||||||
}
|
|
||||||
|
|
||||||
func guardedFiles(t *testing.T) map[string][]string {
|
func guardedFiles(t *testing.T) map[string][]string {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|||||||
@@ -0,0 +1,83 @@
|
|||||||
|
package resolvemcp
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/rand"
|
||||||
|
"crypto/sha256"
|
||||||
|
"encoding/hex"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// confirmStore holds the single-use confirmation tokens for filter-based writes. It is
|
||||||
|
// in-memory: tokens are lost on restart and are not shared between instances, which only costs
|
||||||
|
// the client one more preview call.
|
||||||
|
type confirmStore struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
tokens map[string]confirmEntry
|
||||||
|
now func() time.Time
|
||||||
|
maxLive int
|
||||||
|
}
|
||||||
|
|
||||||
|
type confirmEntry struct {
|
||||||
|
user, table, op, binding string
|
||||||
|
expires time.Time
|
||||||
|
}
|
||||||
|
|
||||||
|
func newConfirmStore() *confirmStore {
|
||||||
|
return &confirmStore{tokens: map[string]confirmEntry{}, now: time.Now, maxLive: 10000}
|
||||||
|
}
|
||||||
|
|
||||||
|
var errConfirmInvalid = NewClientError(CodeInvalidArgument, "confirm_token is invalid or expired; repeat the call without it to get a new preview")
|
||||||
|
|
||||||
|
// issue returns a token bound to the caller, table, operation and binding (a hash of the
|
||||||
|
// filters, data and matched rows the preview showed).
|
||||||
|
func (c *confirmStore) issue(user, table, op, binding string, ttl time.Duration) (string, error) {
|
||||||
|
var b [16]byte
|
||||||
|
if _, err := rand.Read(b[:]); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
tok := hex.EncodeToString(b[:])
|
||||||
|
c.mu.Lock()
|
||||||
|
defer c.mu.Unlock()
|
||||||
|
now := c.now()
|
||||||
|
for k, e := range c.tokens {
|
||||||
|
if now.After(e.expires) {
|
||||||
|
delete(c.tokens, k)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(c.tokens) >= c.maxLive {
|
||||||
|
return "", NewClientError(CodeLimitExceeded, "too many pending confirmations; try again later")
|
||||||
|
}
|
||||||
|
c.tokens[tok] = confirmEntry{user: user, table: table, op: op, binding: binding, expires: now.Add(ttl)}
|
||||||
|
return tok, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// consume validates and removes a token. Any mismatch (other user, table, operation or
|
||||||
|
// changed binding) is the same error, and the token is spent either way.
|
||||||
|
func (c *confirmStore) consume(tok, user, table, op, binding string) error {
|
||||||
|
c.mu.Lock()
|
||||||
|
e, ok := c.tokens[tok]
|
||||||
|
delete(c.tokens, tok)
|
||||||
|
now := c.now()
|
||||||
|
c.mu.Unlock()
|
||||||
|
if !ok || now.After(e.expires) || e.user != user || e.table != table || e.op != op || e.binding != binding {
|
||||||
|
return errConfirmInvalid
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// bindingHash fingerprints the parts of a write a confirmation covers.
|
||||||
|
func bindingHash(parts ...any) (string, error) {
|
||||||
|
h := sha256.New()
|
||||||
|
for _, p := range parts {
|
||||||
|
b, err := json.Marshal(p)
|
||||||
|
if err != nil {
|
||||||
|
return "", errors.New("cannot fingerprint request")
|
||||||
|
}
|
||||||
|
h.Write(b)
|
||||||
|
h.Write([]byte{0})
|
||||||
|
}
|
||||||
|
return hex.EncodeToString(h.Sum(nil)), nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,280 @@
|
|||||||
|
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
|
||||||
|
}
|
||||||
@@ -33,6 +33,8 @@ type Handler struct {
|
|||||||
version string
|
version string
|
||||||
oauth2Regs []oauth2Registration
|
oauth2Regs []oauth2Registration
|
||||||
oauthSrv *security.OAuthServer
|
oauthSrv *security.OAuthServer
|
||||||
|
functions functionRegistry
|
||||||
|
confirms *confirmStore
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewHandler creates a Handler with the given database, model registry, and config.
|
// NewHandler creates a Handler with the given database, model registry, and config.
|
||||||
@@ -43,9 +45,11 @@ func NewHandler(db common.Database, registry common.ModelRegistry, cfg Config) *
|
|||||||
hooks: NewHookRegistry(),
|
hooks: NewHookRegistry(),
|
||||||
mcpServer: server.NewMCPServer("resolvemcp", "1.0.0"),
|
mcpServer: server.NewMCPServer("resolvemcp", "1.0.0"),
|
||||||
config: cfg.withDefaults(),
|
config: cfg.withDefaults(),
|
||||||
|
confirms: newConfirmStore(),
|
||||||
name: "resolvemcp",
|
name: "resolvemcp",
|
||||||
version: "1.0.0",
|
version: "1.0.0",
|
||||||
}
|
}
|
||||||
|
registerMetaTools(h)
|
||||||
if cfg.EnableAnnotations {
|
if cfg.EnableAnnotations {
|
||||||
registerAnnotationTool(h)
|
registerAnnotationTool(h)
|
||||||
}
|
}
|
||||||
@@ -165,13 +169,13 @@ func requestBaseURL(r *http.Request) string {
|
|||||||
return scheme + "://" + r.Host
|
return scheme + "://" + r.Host
|
||||||
}
|
}
|
||||||
|
|
||||||
// RegisterModel registers a model and immediately exposes it as MCP tools and a resource.
|
// RegisterModel registers a model. It becomes visible to the fixed meta tools (list_tables,
|
||||||
|
// select_table, ...); no per-model tools are created.
|
||||||
func (h *Handler) RegisterModel(schema, entity string, model interface{}) error {
|
func (h *Handler) RegisterModel(schema, entity string, model interface{}) error {
|
||||||
fullName := buildModelName(schema, entity)
|
fullName := buildModelName(schema, entity)
|
||||||
if err := h.registry.RegisterModel(fullName, model); err != nil {
|
if err := h.registry.RegisterModel(fullName, model); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
registerModelTools(h, schema, entity, model)
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -187,7 +191,6 @@ func (h *Handler) RegisterModelWithRules(schema, entity string, model interface{
|
|||||||
if err := reg.RegisterModelWithRules(fullName, model, rules); err != nil {
|
if err := reg.RegisterModelWithRules(fullName, model, rules); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
registerModelTools(h, schema, entity, model)
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -34,6 +34,12 @@ const (
|
|||||||
BeforeDelete HookType = "before_delete"
|
BeforeDelete HookType = "before_delete"
|
||||||
AfterDelete HookType = "after_delete"
|
AfterDelete HookType = "after_delete"
|
||||||
|
|
||||||
|
// BeforeCall and AfterCall fire inside the transaction of a call_function call.
|
||||||
|
// hookCtx.Entity is the function name, Data the validated arguments (BeforeCall may
|
||||||
|
// replace them) and Result the function's result (AfterCall).
|
||||||
|
BeforeCall HookType = "before_call"
|
||||||
|
AfterCall HookType = "after_call"
|
||||||
|
|
||||||
// OnTxBegin fires once, first, inside every transaction the handler opens
|
// OnTxBegin fires once, first, inside every transaction the handler opens
|
||||||
// (including the second short transaction for post-commit work). hookCtx.Tx is
|
// (including the second short transaction for post-commit work). hookCtx.Tx is
|
||||||
// the transaction; use it to stamp transaction-local state such as RLS
|
// the transaction; use it to stamp transaction-local state such as RLS
|
||||||
|
|||||||
@@ -0,0 +1,438 @@
|
|||||||
|
package resolvemcp
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"reflect"
|
||||||
|
"sort"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/mark3labs/mcp-go/mcp"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Operation names used by list_tables and the rule checks.
|
||||||
|
const (
|
||||||
|
opSelect = "select"
|
||||||
|
opInsert = "insert"
|
||||||
|
opUpdate = "update"
|
||||||
|
opDelete = "delete"
|
||||||
|
)
|
||||||
|
|
||||||
|
// filterOperators is the operator list shown by describe_table.
|
||||||
|
var filterOperators = []string{"=", "!=", ">", ">=", "<", "<=", "like", "ilike", "in", "is_null", "is_not_null"}
|
||||||
|
|
||||||
|
// registerMetaTools adds the fixed tool set. Their number does not grow with the models.
|
||||||
|
func registerMetaTools(h *Handler) {
|
||||||
|
readOnly := mcp.WithReadOnlyHintAnnotation(true)
|
||||||
|
tableArg := mcp.WithString("table", mcp.Required(), mcp.Description("Table as 'schema.entity' (see list_tables)."))
|
||||||
|
filtersArg := mcp.WithArray("filters", mcp.Description(`Filter objects, e.g. [{"column":"status","operator":"=","value":"active"}]. Combine with "logic_operator": "AND" (default) or "OR". Operators: `+strings.Join(filterOperators, " ")+"."))
|
||||||
|
idArg := mcp.WithString("id", mcp.Description("Primary key of one row. Use either id or filters."))
|
||||||
|
dryRunArg := mcp.WithBoolean("dry_run", mcp.Description("Report how many rows match and a preview of their ids; change nothing."))
|
||||||
|
confirmArg := mcp.WithString("confirm_token", mcp.Description("Token from the preview a filter-based call returns first; the write only happens with it."))
|
||||||
|
|
||||||
|
h.mcpServer.AddTool(mcp.NewTool("list_tables", readOnly,
|
||||||
|
mcp.WithDescription("List the tables you can use and the operations (select, insert, update, delete) allowed on each.")),
|
||||||
|
h.handleListTables)
|
||||||
|
|
||||||
|
h.mcpServer.AddTool(mcp.NewTool("describe_table", readOnly,
|
||||||
|
mcp.WithDescription("Describe a table: columns and types, primary key, relations (preloadable), writable fields, allowed operations and server limits."),
|
||||||
|
tableArg), h.handleDescribeTable)
|
||||||
|
|
||||||
|
h.mcpServer.AddTool(mcp.NewTool("select_table", readOnly,
|
||||||
|
mcp.WithDescription("Read rows from a table. Results are paged: 'limit' defaults to the server default and is capped; use cursor_forward/cursor_backward for deep paging. The total row count is only computed with include_count."),
|
||||||
|
tableArg, idArg, filtersArg,
|
||||||
|
mcp.WithArray("sort", mcp.Description(`Sort objects, e.g. [{"column":"created_at","direction":"desc"}].`)),
|
||||||
|
mcp.WithArray("columns", mcp.Description("Columns to return. Omit for all.")),
|
||||||
|
mcp.WithArray("omit_columns", mcp.Description("Columns to leave out.")),
|
||||||
|
mcp.WithArray("preloads", mcp.Description(`Relations to load, e.g. [{"relation":"orders"}]. See describe_table for the names and the maximum depth.`)),
|
||||||
|
mcp.WithNumber("limit", mcp.Description("Maximum rows to return.")),
|
||||||
|
mcp.WithNumber("offset", mcp.Description("Rows to skip.")),
|
||||||
|
mcp.WithString("cursor_forward", mcp.Description("Primary key of the last row of the current page; requires sort.")),
|
||||||
|
mcp.WithString("cursor_backward", mcp.Description("Primary key of the first row of the current page; requires sort.")),
|
||||||
|
mcp.WithBoolean("include_count", mcp.Description("Also return the total number of matching rows (slower on large tables).")),
|
||||||
|
), h.handleSelect)
|
||||||
|
|
||||||
|
h.mcpServer.AddTool(mcp.NewTool("insert_into_table",
|
||||||
|
mcp.WithDescription("Insert one row (object) or several rows (array, one transaction, capped). Unknown or read-only fields are rejected."),
|
||||||
|
tableArg, mcp.WithObject("data", mcp.Required(), mcp.Description("A row object or an array of row objects.")),
|
||||||
|
), h.handleInsert)
|
||||||
|
|
||||||
|
h.mcpServer.AddTool(mcp.NewTool("update_table",
|
||||||
|
mcp.WithDescription("Update rows. Give an id (one row, applied at once) or filters (several rows: the first call returns a preview and a confirm_token, repeat the call with the token to apply; the number of rows is capped). Only the fields in data are changed; null sets NULL."),
|
||||||
|
tableArg, idArg, filtersArg,
|
||||||
|
mcp.WithObject("data", mcp.Required(), mcp.Description("Fields to change.")),
|
||||||
|
dryRunArg, confirmArg,
|
||||||
|
), h.handleUpdate)
|
||||||
|
|
||||||
|
h.mcpServer.AddTool(mcp.NewTool("delete_from_table",
|
||||||
|
mcp.WithDescription("Delete rows. Give an id (one row, applied at once) or filters (several rows: the first call returns a preview and a confirm_token, repeat the call with the token to delete; the number of rows is capped)."),
|
||||||
|
mcp.WithDestructiveHintAnnotation(true),
|
||||||
|
tableArg, idArg, filtersArg, dryRunArg, confirmArg,
|
||||||
|
), h.handleDelete)
|
||||||
|
|
||||||
|
h.mcpServer.AddTool(mcp.NewTool("list_functions", readOnly,
|
||||||
|
mcp.WithDescription("List the functions you can call with call_function, with their parameters.")),
|
||||||
|
h.handleListFunctions)
|
||||||
|
|
||||||
|
h.mcpServer.AddTool(mcp.NewTool("call_function",
|
||||||
|
mcp.WithDescription("Call a registered function by name with validated arguments (see list_functions). Runs in a transaction."),
|
||||||
|
mcp.WithString("name", mcp.Required(), mcp.Description("Function name.")),
|
||||||
|
mcp.WithObject("arguments", mcp.Description("Arguments by parameter name.")),
|
||||||
|
), h.handleCallFunction)
|
||||||
|
}
|
||||||
|
|
||||||
|
// --------------------------------------------------------------------------
|
||||||
|
// Tables and rules
|
||||||
|
// --------------------------------------------------------------------------
|
||||||
|
|
||||||
|
// splitTable parses "schema.entity" (or a bare entity).
|
||||||
|
func splitTable(table string) (schema, entity string, err error) {
|
||||||
|
table = strings.TrimSpace(table)
|
||||||
|
if table == "" {
|
||||||
|
return "", "", invalidArg("missing required argument: table")
|
||||||
|
}
|
||||||
|
if i := strings.LastIndex(table, "."); i >= 0 {
|
||||||
|
return table[:i], table[i+1:], nil
|
||||||
|
}
|
||||||
|
return "", table, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// modelRules returns the rules of a registered model (defaults when the registry keeps none).
|
||||||
|
func (h *Handler) modelRules(schema, entity string) modelregistry.ModelRules {
|
||||||
|
if reg, ok := h.registry.(*modelregistry.DefaultModelRegistry); ok {
|
||||||
|
if r, err := reg.GetModelRules(buildModelName(schema, entity)); err == nil {
|
||||||
|
return r
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return modelregistry.DefaultModelRules()
|
||||||
|
}
|
||||||
|
|
||||||
|
func opsFor(r modelregistry.ModelRules) []string {
|
||||||
|
var ops []string
|
||||||
|
if r.CanRead {
|
||||||
|
ops = append(ops, opSelect)
|
||||||
|
}
|
||||||
|
if r.CanCreate {
|
||||||
|
ops = append(ops, opInsert)
|
||||||
|
}
|
||||||
|
if r.CanUpdate {
|
||||||
|
ops = append(ops, opUpdate)
|
||||||
|
}
|
||||||
|
if r.CanDelete {
|
||||||
|
ops = append(ops, opDelete)
|
||||||
|
}
|
||||||
|
return ops
|
||||||
|
}
|
||||||
|
|
||||||
|
// resolveTable finds a registered model and checks the rule for op. A model the rules forbid
|
||||||
|
// for op is reported like a missing one for reads of the table list, but with a plain
|
||||||
|
// "not allowed" here so the agent learns the operation is off, not that it misspelled the name.
|
||||||
|
func (h *Handler) resolveTable(args map[string]any, op string) (schema, entity string, err error) {
|
||||||
|
table, _ := args["table"].(string)
|
||||||
|
schema, entity, err = splitTable(table)
|
||||||
|
if err != nil {
|
||||||
|
return "", "", err
|
||||||
|
}
|
||||||
|
if _, err := h.registry.GetModelByEntity(schema, entity); err != nil {
|
||||||
|
return "", "", invalidArg("unknown table %q; see list_tables", truncate(table))
|
||||||
|
}
|
||||||
|
if op != "" {
|
||||||
|
allowed := false
|
||||||
|
for _, o := range opsFor(h.modelRules(schema, entity)) {
|
||||||
|
if o == op {
|
||||||
|
allowed = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !allowed {
|
||||||
|
return "", "", NewClientError(CodeForbidden, fmt.Sprintf("%s is not allowed on %s", op, buildModelName(schema, entity)))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return schema, entity, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) handleListTables(ctx context.Context, _ mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||||
|
type table struct {
|
||||||
|
Table string `json:"table"`
|
||||||
|
Operations []string `json:"operations"`
|
||||||
|
}
|
||||||
|
var tables []table
|
||||||
|
for name := range h.registry.GetAllModels() {
|
||||||
|
schema, entity, _ := splitTable(name)
|
||||||
|
if ops := opsFor(h.modelRules(schema, entity)); len(ops) > 0 {
|
||||||
|
tables = append(tables, table{Table: name, Operations: ops})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
sort.Slice(tables, func(i, j int) bool { return tables[i].Table < tables[j].Table })
|
||||||
|
return marshalResult(map[string]any{"success": true, "tables": tables})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) handleDescribeTable(_ context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||||
|
schema, entity, err := h.resolveTable(req.GetArguments(), "")
|
||||||
|
if err != nil {
|
||||||
|
return toolError("describe_table", err), nil
|
||||||
|
}
|
||||||
|
model, err := h.registry.GetModelByEntity(schema, entity)
|
||||||
|
if err != nil {
|
||||||
|
return toolError("describe_table", invalidArg("unknown table")), nil
|
||||||
|
}
|
||||||
|
rules := h.modelRules(schema, entity)
|
||||||
|
if len(opsFor(rules)) == 0 {
|
||||||
|
return toolError("describe_table", invalidArg("unknown table %q; see list_tables", buildModelName(schema, entity))), nil
|
||||||
|
}
|
||||||
|
info := buildModelInfo(schema, entity, model)
|
||||||
|
|
||||||
|
modelType := reflect.TypeOf(model)
|
||||||
|
for modelType != nil && (modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice) {
|
||||||
|
modelType = modelType.Elem()
|
||||||
|
}
|
||||||
|
writable := map[string]bool{}
|
||||||
|
if modelType != nil && modelType.Kind() == reflect.Struct {
|
||||||
|
for jsonKey := range reflection.BuildJSONToDBColumnMap(modelType) {
|
||||||
|
writable[jsonKey] = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type column struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Type string `json:"type,omitempty"`
|
||||||
|
Nullable bool `json:"nullable"`
|
||||||
|
PrimaryKey bool `json:"primary_key,omitempty"`
|
||||||
|
Unique bool `json:"unique,omitempty"`
|
||||||
|
Writable bool `json:"writable"`
|
||||||
|
}
|
||||||
|
cols := make([]column, 0, len(info.columns))
|
||||||
|
var writableNames []string
|
||||||
|
for _, c := range info.columns {
|
||||||
|
typ := c.sqlType
|
||||||
|
if typ == "" {
|
||||||
|
typ = c.goType
|
||||||
|
}
|
||||||
|
w := writable[c.jsonName]
|
||||||
|
cols = append(cols, column{Name: c.jsonName, Type: typ, Nullable: c.nullable, PrimaryKey: c.isPrimary, Unique: c.isUnique, Writable: w})
|
||||||
|
if w && !c.isPrimary {
|
||||||
|
writableNames = append(writableNames, c.jsonName)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return marshalResult(map[string]any{
|
||||||
|
"success": true,
|
||||||
|
"table": info.fullName,
|
||||||
|
"primary_key": info.pkName,
|
||||||
|
"columns": cols,
|
||||||
|
"relations": info.relationNames,
|
||||||
|
"writable_columns": writableNames,
|
||||||
|
"operations": opsFor(rules),
|
||||||
|
"filter_operators": filterOperators,
|
||||||
|
"limits": map[string]any{
|
||||||
|
"default_limit": h.config.DefaultLimit,
|
||||||
|
"max_limit": h.config.MaxLimit,
|
||||||
|
"max_offset": h.config.MaxOffset,
|
||||||
|
"max_batch": h.config.MaxBatch,
|
||||||
|
"max_preload_depth": h.config.MaxPreloadDepth,
|
||||||
|
"max_write_rows": h.config.MaxWriteRows,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// --------------------------------------------------------------------------
|
||||||
|
// Row tools
|
||||||
|
// --------------------------------------------------------------------------
|
||||||
|
|
||||||
|
// argID reads the id argument, which clients send as a string or a number.
|
||||||
|
func argID(args map[string]any) string {
|
||||||
|
switch v := args["id"].(type) {
|
||||||
|
case string:
|
||||||
|
return v
|
||||||
|
case float64:
|
||||||
|
if v == float64(int64(v)) {
|
||||||
|
return fmt.Sprintf("%d", int64(v))
|
||||||
|
}
|
||||||
|
return fmt.Sprint(v)
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseFiltersStrict is parseFilters for writes: a malformed filter is an error, never dropped.
|
||||||
|
func parseFiltersStrict(raw any) ([]common.FilterOption, error) {
|
||||||
|
if raw == nil {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
items, ok := raw.([]any)
|
||||||
|
if !ok {
|
||||||
|
return nil, invalidArg("filters must be an array")
|
||||||
|
}
|
||||||
|
parsed := parseFilters(raw)
|
||||||
|
if len(parsed) != len(items) {
|
||||||
|
return nil, invalidArg("every filter needs a column and an operator")
|
||||||
|
}
|
||||||
|
return parsed, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) handleSelect(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||||
|
args := req.GetArguments()
|
||||||
|
schema, entity, err := h.resolveTable(args, opSelect)
|
||||||
|
if err != nil {
|
||||||
|
return toolError("select_table", err), nil
|
||||||
|
}
|
||||||
|
count, _ := args["include_count"].(bool)
|
||||||
|
data, meta, err := h.executeReadCounted(ctx, schema, entity, argID(args), parseRequestOptions(args), count)
|
||||||
|
if err != nil {
|
||||||
|
return toolError("select_table", err), nil
|
||||||
|
}
|
||||||
|
if !count && meta != nil {
|
||||||
|
meta.Total, meta.Filtered = 0, 0
|
||||||
|
}
|
||||||
|
return marshalResult(map[string]any{"success": true, "data": data, "metadata": meta})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) handleInsert(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||||
|
args := req.GetArguments()
|
||||||
|
schema, entity, err := h.resolveTable(args, opInsert)
|
||||||
|
if err != nil {
|
||||||
|
return toolError("insert_into_table", err), nil
|
||||||
|
}
|
||||||
|
data, ok := args["data"]
|
||||||
|
if !ok {
|
||||||
|
return toolError("insert_into_table", invalidArg("missing required argument: data")), nil
|
||||||
|
}
|
||||||
|
result, err := h.executeCreate(ctx, schema, entity, data)
|
||||||
|
if err != nil {
|
||||||
|
return toolError("insert_into_table", err), nil
|
||||||
|
}
|
||||||
|
return marshalResult(map[string]any{"success": true, "data": result})
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeTarget parses id/filters/dry_run/confirm_token shared by update and delete.
|
||||||
|
func writeTarget(args map[string]any) (id string, filters []common.FilterOption, dryRun bool, token string, err error) {
|
||||||
|
id = argID(args)
|
||||||
|
if filters, err = parseFiltersStrict(args["filters"]); err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if id != "" && len(filters) > 0 {
|
||||||
|
err = invalidArg("use either id or filters, not both")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if id == "" && len(filters) == 0 {
|
||||||
|
err = invalidArg("provide an id or at least one filter")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
dryRun, _ = args["dry_run"].(bool)
|
||||||
|
token, _ = args["confirm_token"].(string)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) handleUpdate(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||||
|
args := req.GetArguments()
|
||||||
|
schema, entity, err := h.resolveTable(args, opUpdate)
|
||||||
|
if err != nil {
|
||||||
|
return toolError("update_table", err), nil
|
||||||
|
}
|
||||||
|
data, ok := args["data"].(map[string]any)
|
||||||
|
if !ok {
|
||||||
|
return toolError("update_table", invalidArg("data must be an object")), nil
|
||||||
|
}
|
||||||
|
id, filters, dryRun, token, err := writeTarget(args)
|
||||||
|
if err != nil {
|
||||||
|
return toolError("update_table", err), nil
|
||||||
|
}
|
||||||
|
if id != "" && !dryRun {
|
||||||
|
result, err := h.executeUpdate(ctx, schema, entity, id, data)
|
||||||
|
if err != nil {
|
||||||
|
return toolError("update_table", err), nil
|
||||||
|
}
|
||||||
|
return marshalResult(map[string]any{"success": true, "data": result})
|
||||||
|
}
|
||||||
|
return h.runWhere(ctx, "update_table", whereRequest{schema: schema, entity: entity, op: "update", filters: h.idFilters(schema, entity, id, filters), data: data, dryRun: dryRun, confirmToken: token})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) handleDelete(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||||
|
args := req.GetArguments()
|
||||||
|
schema, entity, err := h.resolveTable(args, opDelete)
|
||||||
|
if err != nil {
|
||||||
|
return toolError("delete_from_table", err), nil
|
||||||
|
}
|
||||||
|
id, filters, dryRun, token, err := writeTarget(args)
|
||||||
|
if err != nil {
|
||||||
|
return toolError("delete_from_table", err), nil
|
||||||
|
}
|
||||||
|
if id != "" && !dryRun {
|
||||||
|
result, err := h.executeDelete(ctx, schema, entity, id)
|
||||||
|
if err != nil {
|
||||||
|
return toolError("delete_from_table", err), nil
|
||||||
|
}
|
||||||
|
return marshalResult(map[string]any{"success": true, "data": result})
|
||||||
|
}
|
||||||
|
return h.runWhere(ctx, "delete_from_table", whereRequest{schema: schema, entity: entity, op: "delete", filters: h.idFilters(schema, entity, id, filters), dryRun: dryRun, confirmToken: token})
|
||||||
|
}
|
||||||
|
|
||||||
|
// idFilters turns an id into a primary-key filter so a dry run of an id write goes through the
|
||||||
|
// same matching as a filter write.
|
||||||
|
func (h *Handler) idFilters(schema, entity, id string, filters []common.FilterOption) []common.FilterOption {
|
||||||
|
if id == "" {
|
||||||
|
return filters
|
||||||
|
}
|
||||||
|
model, err := h.registry.GetModelByEntity(schema, entity)
|
||||||
|
if err != nil {
|
||||||
|
return filters
|
||||||
|
}
|
||||||
|
return []common.FilterOption{{Column: reflection.GetPrimaryKeyName(model), Operator: "eq", Value: id, LogicOperator: "AND"}}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) runWhere(ctx context.Context, op string, req whereRequest) (*mcp.CallToolResult, error) {
|
||||||
|
res, err := h.executeWhere(ctx, req)
|
||||||
|
if err != nil {
|
||||||
|
return toolError(op, err), nil
|
||||||
|
}
|
||||||
|
return marshalResult(map[string]any{"success": true, "result": res})
|
||||||
|
}
|
||||||
|
|
||||||
|
// --------------------------------------------------------------------------
|
||||||
|
// Functions
|
||||||
|
// --------------------------------------------------------------------------
|
||||||
|
|
||||||
|
func (h *Handler) handleListFunctions(ctx context.Context, _ mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||||
|
type param struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
Required bool `json:"required"`
|
||||||
|
Description string `json:"description,omitempty"`
|
||||||
|
}
|
||||||
|
type fn struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Description string `json:"description,omitempty"`
|
||||||
|
Parameters []param `json:"parameters"`
|
||||||
|
}
|
||||||
|
out := []fn{}
|
||||||
|
for _, f := range h.visibleFunctions(ctx) {
|
||||||
|
item := fn{Name: f.Name, Description: f.Description, Parameters: []param{}}
|
||||||
|
for _, p := range f.Params {
|
||||||
|
item.Parameters = append(item.Parameters, param{p.Name, p.Type, p.Required, p.Description})
|
||||||
|
}
|
||||||
|
out = append(out, item)
|
||||||
|
}
|
||||||
|
return marshalResult(map[string]any{"success": true, "functions": out})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) handleCallFunction(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||||
|
args := req.GetArguments()
|
||||||
|
name, _ := args["name"].(string)
|
||||||
|
if name == "" {
|
||||||
|
return toolError("call_function", invalidArg("missing required argument: name")), nil
|
||||||
|
}
|
||||||
|
fnArgs := map[string]any{}
|
||||||
|
if raw, ok := args["arguments"]; ok && raw != nil {
|
||||||
|
m, ok := raw.(map[string]any)
|
||||||
|
if !ok {
|
||||||
|
return toolError("call_function", invalidArg("arguments must be an object")), nil
|
||||||
|
}
|
||||||
|
fnArgs = m
|
||||||
|
}
|
||||||
|
result, err := h.executeCall(ctx, name, fnArgs)
|
||||||
|
if err != nil {
|
||||||
|
return toolError("call_function", err), nil
|
||||||
|
}
|
||||||
|
return marshalResult(map[string]any{"success": true, "result": result})
|
||||||
|
}
|
||||||
@@ -0,0 +1,451 @@
|
|||||||
|
package resolvemcp
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"sort"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/DATA-DOG/go-sqlmock"
|
||||||
|
"github.com/mark3labs/mcp-go/mcp"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security"
|
||||||
|
)
|
||||||
|
|
||||||
|
func callReq(args map[string]any) mcp.CallToolRequest {
|
||||||
|
var r mcp.CallToolRequest
|
||||||
|
r.Params.Arguments = args
|
||||||
|
return r
|
||||||
|
}
|
||||||
|
|
||||||
|
// payload decodes a tool result's text.
|
||||||
|
func payload(t *testing.T, res *mcp.CallToolResult) map[string]any {
|
||||||
|
t.Helper()
|
||||||
|
if len(res.Content) == 0 {
|
||||||
|
t.Fatal("empty result")
|
||||||
|
}
|
||||||
|
tc, ok := res.Content[0].(mcp.TextContent)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("content is %T", res.Content[0])
|
||||||
|
}
|
||||||
|
var m map[string]any
|
||||||
|
if err := json.Unmarshal([]byte(tc.Text), &m); err != nil {
|
||||||
|
t.Fatalf("not JSON: %q", tc.Text)
|
||||||
|
}
|
||||||
|
return m
|
||||||
|
}
|
||||||
|
|
||||||
|
func errCode(t *testing.T, res *mcp.CallToolResult) string {
|
||||||
|
t.Helper()
|
||||||
|
if !res.IsError {
|
||||||
|
t.Fatalf("expected an error result, got %v", payload(t, res))
|
||||||
|
}
|
||||||
|
e, _ := payload(t, res)["error"].(map[string]any)
|
||||||
|
code, _ := e["code"].(string)
|
||||||
|
return code
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMetaToolSetIsFixed(t *testing.T) {
|
||||||
|
h, _, _ := newTxHarness(t)
|
||||||
|
for _, name := range []string{"x1", "x2", "x3"} {
|
||||||
|
if err := h.RegisterModel("public", name, &txItem{}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
var got []string
|
||||||
|
for name := range h.mcpServer.ListTools() {
|
||||||
|
got = append(got, name)
|
||||||
|
}
|
||||||
|
sort.Strings(got)
|
||||||
|
want := "call_function delete_from_table describe_table insert_into_table list_functions list_tables select_table update_table"
|
||||||
|
if strings.Join(got, " ") != want {
|
||||||
|
t.Fatalf("tools = %v\nwant %s", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestListTablesShowsOnlyAllowedOperations(t *testing.T) {
|
||||||
|
h, _, ctx := newTxHarness(t)
|
||||||
|
_ = h.RegisterModelWithRules("public", "ro", &txItem{}, modelregistry.ModelRules{CanRead: true})
|
||||||
|
_ = h.RegisterModelWithRules("public", "hidden", &txItem{}, modelregistry.ModelRules{})
|
||||||
|
res, _ := h.handleListTables(ctx, callReq(nil))
|
||||||
|
tables, _ := payload(t, res)["tables"].([]any)
|
||||||
|
seen := map[string][]any{}
|
||||||
|
for _, tb := range tables {
|
||||||
|
m := tb.(map[string]any)
|
||||||
|
seen[m["table"].(string)] = m["operations"].([]any)
|
||||||
|
}
|
||||||
|
if _, ok := seen["public.hidden"]; ok {
|
||||||
|
t.Error("a table with no allowed operation must not be listed")
|
||||||
|
}
|
||||||
|
if ops := seen["public.ro"]; len(ops) != 1 || ops[0] != "select" {
|
||||||
|
t.Errorf("ro ops = %v", ops)
|
||||||
|
}
|
||||||
|
if ops := seen["public.items"]; len(ops) != 4 {
|
||||||
|
t.Errorf("default rules allow all four, got %v", ops)
|
||||||
|
}
|
||||||
|
// describe_table on a table with no allowed operation reads as unknown.
|
||||||
|
res, _ = h.handleDescribeTable(ctx, callReq(map[string]any{"table": "public.hidden"}))
|
||||||
|
if errCode(t, res) != CodeInvalidArgument {
|
||||||
|
t.Error("describe of a hidden table must look like an unknown table")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDescribeTable(t *testing.T) {
|
||||||
|
h, _, ctx := newTxHarness(t)
|
||||||
|
if err := h.RegisterModel("public", "witems", &wItem{}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
res, _ := h.handleDescribeTable(ctx, callReq(map[string]any{"table": "public.witems"}))
|
||||||
|
p := payload(t, res)
|
||||||
|
if p["primary_key"] != "id" {
|
||||||
|
t.Errorf("pk = %v", p["primary_key"])
|
||||||
|
}
|
||||||
|
w, _ := p["writable_columns"].([]any)
|
||||||
|
got := map[string]bool{}
|
||||||
|
for _, c := range w {
|
||||||
|
got[c.(string)] = true
|
||||||
|
}
|
||||||
|
if !got["name"] || !got["fullName"] || got["id"] || got["owner"] {
|
||||||
|
t.Errorf("writable columns = %v", w)
|
||||||
|
}
|
||||||
|
if lim, _ := p["limits"].(map[string]any); lim["max_limit"] != float64(1000) {
|
||||||
|
t.Errorf("limits = %v", p["limits"])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSelectRespectsOperationRule(t *testing.T) {
|
||||||
|
h, _, ctx := newTxHarness(t)
|
||||||
|
_ = h.RegisterModelWithRules("public", "nowrite", &txItem{}, modelregistry.ModelRules{CanRead: true})
|
||||||
|
for tool, fn := range map[string]func(context.Context, mcp.CallToolRequest) (*mcp.CallToolResult, error){
|
||||||
|
"insert": h.handleInsert, "update": h.handleUpdate, "delete": h.handleDelete,
|
||||||
|
} {
|
||||||
|
res, _ := fn(ctx, callReq(map[string]any{"table": "public.nowrite", "data": map[string]any{"name": "a"}, "id": "1"}))
|
||||||
|
if errCode(t, res) != CodeForbidden {
|
||||||
|
t.Errorf("%s on a read-only table must be forbidden", tool)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
res, _ := h.handleSelect(ctx, callReq(map[string]any{"table": "public.missing"}))
|
||||||
|
if errCode(t, res) != CodeInvalidArgument {
|
||||||
|
t.Error("unknown table must be invalid_argument")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSelectCountOnlyWhenRequested(t *testing.T) {
|
||||||
|
h, mock, ctx := newTxHarness(t)
|
||||||
|
mock.ExpectBegin()
|
||||||
|
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(1, "a"))
|
||||||
|
mock.ExpectCommit()
|
||||||
|
res, _ := h.handleSelect(ctx, callReq(map[string]any{"table": "public.items"}))
|
||||||
|
if res.IsError {
|
||||||
|
t.Fatal(payload(t, res))
|
||||||
|
}
|
||||||
|
if err := mock.ExpectationsWereMet(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
mock.ExpectBegin()
|
||||||
|
mock.ExpectQuery(`SELECT COUNT`).WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(41))
|
||||||
|
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(1, "a"))
|
||||||
|
mock.ExpectCommit()
|
||||||
|
res, _ = h.handleSelect(ctx, callReq(map[string]any{"table": "public.items", "include_count": true}))
|
||||||
|
meta, _ := payload(t, res)["metadata"].(map[string]any)
|
||||||
|
if res.IsError || meta["total"] != float64(41) {
|
||||||
|
t.Fatalf("metadata = %v", meta)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- filter writes ---
|
||||||
|
|
||||||
|
func matchRowsQuery(mock sqlmock.Sqlmock, ids ...int) {
|
||||||
|
rows := sqlmock.NewRows([]string{"id"})
|
||||||
|
for _, id := range ids {
|
||||||
|
rows.AddRow(id)
|
||||||
|
}
|
||||||
|
mock.ExpectQuery(`SELECT`).WillReturnRows(rows)
|
||||||
|
}
|
||||||
|
|
||||||
|
func updReq(filters any, extra map[string]any) mcp.CallToolRequest {
|
||||||
|
a := map[string]any{"table": "public.items", "data": map[string]any{"name": "z"}, "filters": filters}
|
||||||
|
for k, v := range extra {
|
||||||
|
a[k] = v
|
||||||
|
}
|
||||||
|
return callReq(a)
|
||||||
|
}
|
||||||
|
|
||||||
|
var statusFilter = []any{map[string]any{"column": "name", "operator": "=", "value": "a"}}
|
||||||
|
|
||||||
|
func TestFilterUpdateNeedsPreviewThenToken(t *testing.T) {
|
||||||
|
h, mock, ctx := newTxHarness(t)
|
||||||
|
|
||||||
|
// 1. preview: counts, lists ids, issues a token, writes nothing.
|
||||||
|
mock.ExpectBegin()
|
||||||
|
matchRowsQuery(mock, 1, 2)
|
||||||
|
mock.ExpectCommit()
|
||||||
|
res, _ := h.handleUpdate(ctx, updReq(statusFilter, nil))
|
||||||
|
r, _ := payload(t, res)["result"].(map[string]any)
|
||||||
|
tok, _ := r["confirm_token"].(string)
|
||||||
|
if res.IsError || tok == "" || r["requires_confirmation"] != true || r["matched"] != float64(2) {
|
||||||
|
t.Fatalf("preview = %v", payload(t, res))
|
||||||
|
}
|
||||||
|
|
||||||
|
// 2. with the token: re-matches inside the tx, then writes.
|
||||||
|
mock.ExpectBegin()
|
||||||
|
matchRowsQuery(mock, 1, 2)
|
||||||
|
mock.ExpectExec(`UPDATE .* WHERE "id" IN \(\$2, \$3\)`).WillReturnResult(sqlmock.NewResult(0, 2))
|
||||||
|
mock.ExpectCommit()
|
||||||
|
res, _ = h.handleUpdate(ctx, updReq(statusFilter, map[string]any{"confirm_token": tok}))
|
||||||
|
r, _ = payload(t, res)["result"].(map[string]any)
|
||||||
|
if res.IsError || r["affected"] != float64(2) {
|
||||||
|
t.Fatalf("confirmed = %v", payload(t, res))
|
||||||
|
}
|
||||||
|
if err := mock.ExpectationsWereMet(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 3. a token is single use.
|
||||||
|
mock.ExpectBegin()
|
||||||
|
matchRowsQuery(mock, 1, 2)
|
||||||
|
mock.ExpectRollback()
|
||||||
|
res, _ = h.handleUpdate(ctx, updReq(statusFilter, map[string]any{"confirm_token": tok}))
|
||||||
|
if errCode(t, res) != CodeInvalidArgument {
|
||||||
|
t.Error("a spent token must be rejected")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConfirmTokenBinding(t *testing.T) {
|
||||||
|
issue := func(t *testing.T) (*Handler, sqlmock.Sqlmock, context.Context, string) {
|
||||||
|
h, mock, ctx := newTxHarness(t)
|
||||||
|
mock.ExpectBegin()
|
||||||
|
matchRowsQuery(mock, 1)
|
||||||
|
mock.ExpectCommit()
|
||||||
|
res, _ := h.handleUpdate(ctx, updReq(statusFilter, nil))
|
||||||
|
r, _ := payload(t, res)["result"].(map[string]any)
|
||||||
|
return h, mock, ctx, r["confirm_token"].(string)
|
||||||
|
}
|
||||||
|
reject := func(t *testing.T, h *Handler, mock sqlmock.Sqlmock, ctx context.Context, req mcp.CallToolRequest, rows ...int) {
|
||||||
|
t.Helper()
|
||||||
|
mock.ExpectBegin()
|
||||||
|
matchRowsQuery(mock, rows...)
|
||||||
|
mock.ExpectRollback()
|
||||||
|
res, _ := h.handleUpdate(ctx, req)
|
||||||
|
if errCode(t, res) != CodeInvalidArgument {
|
||||||
|
t.Fatal("token must be rejected")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
t.Run("changed data", func(t *testing.T) {
|
||||||
|
h, mock, ctx, tok := issue(t)
|
||||||
|
req := updReq(statusFilter, map[string]any{"confirm_token": tok, "data": map[string]any{"name": "different"}})
|
||||||
|
reject(t, h, mock, ctx, req, 1)
|
||||||
|
})
|
||||||
|
t.Run("changed filters", func(t *testing.T) {
|
||||||
|
h, mock, ctx, tok := issue(t)
|
||||||
|
other := []any{map[string]any{"column": "name", "operator": "=", "value": "b"}}
|
||||||
|
reject(t, h, mock, ctx, updReq(other, map[string]any{"confirm_token": tok}), 1)
|
||||||
|
})
|
||||||
|
t.Run("rows changed since preview", func(t *testing.T) {
|
||||||
|
h, mock, ctx, tok := issue(t)
|
||||||
|
reject(t, h, mock, ctx, updReq(statusFilter, map[string]any{"confirm_token": tok}), 1, 2)
|
||||||
|
})
|
||||||
|
t.Run("other user", func(t *testing.T) {
|
||||||
|
h, mock, ctx, tok := issue(t)
|
||||||
|
other := context.WithValue(ctx, security.UserContextKey, &security.UserContext{UserID: 99, UserName: "mallory"})
|
||||||
|
reject(t, h, mock, other, updReq(statusFilter, map[string]any{"confirm_token": tok}), 1)
|
||||||
|
})
|
||||||
|
t.Run("other table", func(t *testing.T) {
|
||||||
|
h, mock, ctx, tok := issue(t)
|
||||||
|
if err := h.RegisterModel("public", "other", &txItem{}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
req := updReq(statusFilter, map[string]any{"confirm_token": tok, "table": "public.other"})
|
||||||
|
reject(t, h, mock, ctx, req, 1)
|
||||||
|
})
|
||||||
|
t.Run("expired", func(t *testing.T) {
|
||||||
|
h, mock, ctx, tok := issue(t)
|
||||||
|
h.confirms.now = func() time.Time { return time.Now().Add(time.Hour) }
|
||||||
|
reject(t, h, mock, ctx, updReq(statusFilter, map[string]any{"confirm_token": tok}), 1)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFilterWriteGuardrails(t *testing.T) {
|
||||||
|
h, mock, ctx := newTxHarness(t)
|
||||||
|
for name, req := range map[string]mcp.CallToolRequest{
|
||||||
|
"neither id nor filters": callReq(map[string]any{"table": "public.items", "data": map[string]any{"name": "z"}}),
|
||||||
|
"both id and filters": updReq(statusFilter, map[string]any{"id": "1"}),
|
||||||
|
"malformed filter": updReq([]any{map[string]any{"column": "name"}}, nil),
|
||||||
|
"unknown column": updReq([]any{map[string]any{"column": "secret", "operator": "=", "value": 1}}, nil),
|
||||||
|
"injection in column": updReq([]any{map[string]any{"column": "name) OR (1=1", "operator": "=", "value": 1}}, nil),
|
||||||
|
"unknown operator": updReq([]any{map[string]any{"column": "name", "operator": "ɸ", "value": 1}}, nil),
|
||||||
|
"missing value": updReq([]any{map[string]any{"column": "name", "operator": "="}}, nil),
|
||||||
|
"unknown data field": callReq(map[string]any{"table": "public.items", "data": map[string]any{"role": "x"}, "filters": statusFilter}),
|
||||||
|
} {
|
||||||
|
res, _ := h.handleUpdate(ctx, req)
|
||||||
|
if errCode(t, res) != CodeInvalidArgument {
|
||||||
|
t.Errorf("%s: want invalid_argument", name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Rejected before any SQL ran.
|
||||||
|
if err := mock.ExpectationsWereMet(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFilterWriteRowCap(t *testing.T) {
|
||||||
|
h, mock, ctx := newTxHarness(t)
|
||||||
|
h.config.MaxWriteRows = 2
|
||||||
|
mock.ExpectBegin()
|
||||||
|
matchRowsQuery(mock, 1, 2, 3)
|
||||||
|
mock.ExpectRollback()
|
||||||
|
res, _ := h.handleDelete(ctx, callReq(map[string]any{"table": "public.items", "filters": statusFilter}))
|
||||||
|
if errCode(t, res) != CodeLimitExceeded {
|
||||||
|
t.Fatal("want limit_exceeded")
|
||||||
|
}
|
||||||
|
if err := mock.ExpectationsWereMet(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDryRunWritesNothing(t *testing.T) {
|
||||||
|
h, mock, ctx := newTxHarness(t)
|
||||||
|
mock.ExpectBegin()
|
||||||
|
matchRowsQuery(mock, 5)
|
||||||
|
mock.ExpectCommit()
|
||||||
|
res, _ := h.handleDelete(ctx, callReq(map[string]any{"table": "public.items", "filters": statusFilter, "dry_run": true}))
|
||||||
|
r, _ := payload(t, res)["result"].(map[string]any)
|
||||||
|
if res.IsError || r["dry_run"] != true || r["matched"] != float64(1) || r["confirm_token"] != nil {
|
||||||
|
t.Fatalf("dry run = %v", payload(t, res))
|
||||||
|
}
|
||||||
|
if err := mock.ExpectationsWereMet(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIDWriteNeedsNoToken(t *testing.T) {
|
||||||
|
h, mock, ctx := newTxHarness(t)
|
||||||
|
mock.ExpectBegin()
|
||||||
|
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a"))
|
||||||
|
mock.ExpectExec(`DELETE`).WillReturnResult(sqlmock.NewResult(0, 1))
|
||||||
|
mock.ExpectCommit()
|
||||||
|
res, _ := h.handleDelete(ctx, callReq(map[string]any{"table": "public.items", "id": float64(7)}))
|
||||||
|
if res.IsError {
|
||||||
|
t.Fatal(payload(t, res))
|
||||||
|
}
|
||||||
|
if err := mock.ExpectationsWereMet(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- functions ---
|
||||||
|
|
||||||
|
func TestRegisterFunctionValidation(t *testing.T) {
|
||||||
|
h, _, _ := newTxHarness(t)
|
||||||
|
noop := func(context.Context, common.Database, map[string]any) (any, error) { return nil, nil }
|
||||||
|
bad := map[string]Function{
|
||||||
|
"bad name": {Name: "1x", Handler: noop},
|
||||||
|
"neither": {Name: "f"},
|
||||||
|
"both": {Name: "f", Handler: noop, Procedure: "p"},
|
||||||
|
"bad procedure": {Name: "f", Procedure: "p(); drop table x"},
|
||||||
|
"bad param type": {Name: "f", Handler: noop, Params: []FunctionParam{{Name: "a", Type: "blob"}}},
|
||||||
|
"duplicate param": {Name: "f", Handler: noop, Params: []FunctionParam{{Name: "a", Type: "string"}, {Name: "a", Type: "string"}}},
|
||||||
|
"bad param name": {Name: "f", Handler: noop, Params: []FunctionParam{{Name: "a b", Type: "string"}}},
|
||||||
|
}
|
||||||
|
for name, f := range bad {
|
||||||
|
if err := h.RegisterFunction(f); err == nil {
|
||||||
|
t.Errorf("%s: expected error", name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := h.RegisterFunction(Function{Name: "ok", Handler: noop}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := h.RegisterFunction(Function{Name: "ok", Handler: noop}); err == nil {
|
||||||
|
t.Error("duplicate name must fail")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCallFunctionValidatesAndRunsInTx(t *testing.T) {
|
||||||
|
h, mock, ctx := newTxHarness(t)
|
||||||
|
var gotArgs map[string]any
|
||||||
|
var gotTx common.Database
|
||||||
|
if err := h.RegisterFunction(Function{
|
||||||
|
Name: "greet", Description: "says hi",
|
||||||
|
Params: []FunctionParam{{Name: "who", Type: ParamString, Required: true}, {Name: "n", Type: ParamInteger}},
|
||||||
|
Handler: func(_ context.Context, tx common.Database, args map[string]any) (any, error) {
|
||||||
|
gotArgs, gotTx = args, tx
|
||||||
|
return map[string]any{"hello": args["who"]}, nil
|
||||||
|
},
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
tr := traceHooks(h, OnTxBegin, BeforeCall, AfterCall)
|
||||||
|
|
||||||
|
for name, args := range map[string]map[string]any{
|
||||||
|
"missing required": {},
|
||||||
|
"wrong type": {"who": 5},
|
||||||
|
"fractional int": {"who": "x", "n": 1.5},
|
||||||
|
"unknown arg": {"who": "x", "extra": 1},
|
||||||
|
} {
|
||||||
|
res, _ := h.handleCallFunction(ctx, callReq(map[string]any{"name": "greet", "arguments": args}))
|
||||||
|
if errCode(t, res) != CodeInvalidArgument {
|
||||||
|
t.Errorf("%s: want invalid_argument", name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
res, _ := h.handleCallFunction(ctx, callReq(map[string]any{"name": "nope"}))
|
||||||
|
if errCode(t, res) != CodeInvalidArgument {
|
||||||
|
t.Error("unknown function must be invalid_argument")
|
||||||
|
}
|
||||||
|
|
||||||
|
mock.ExpectBegin()
|
||||||
|
mock.ExpectCommit()
|
||||||
|
res, _ = h.handleCallFunction(ctx, callReq(map[string]any{"name": "greet", "arguments": map[string]any{"who": "kim", "n": float64(2)}}))
|
||||||
|
if res.IsError || gotArgs["who"] != "kim" || gotTx == nil {
|
||||||
|
t.Fatalf("call = %v", payload(t, res))
|
||||||
|
}
|
||||||
|
tr.assertOrder(t, "on_tx_begin", "before_call", "after_call")
|
||||||
|
if tr.txs["before_call"][0] != tr.txs["on_tx_begin"][0] {
|
||||||
|
t.Error("the call must run in the OnTxBegin transaction")
|
||||||
|
}
|
||||||
|
if err := mock.ExpectationsWereMet(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFunctionAuthorizeHidesAndBlocks(t *testing.T) {
|
||||||
|
h, _, ctx := newTxHarness(t)
|
||||||
|
noop := func(context.Context, common.Database, map[string]any) (any, error) { return "ran", nil }
|
||||||
|
_ = h.RegisterFunction(Function{Name: "open", Handler: noop})
|
||||||
|
_ = h.RegisterFunction(Function{Name: "admin_only", Handler: noop, Authorize: func(context.Context) error { return errors.New("no") }})
|
||||||
|
|
||||||
|
res, _ := h.handleListFunctions(ctx, callReq(nil))
|
||||||
|
fns, _ := payload(t, res)["functions"].([]any)
|
||||||
|
if len(fns) != 1 || fns[0].(map[string]any)["name"] != "open" {
|
||||||
|
t.Fatalf("visible functions = %v", fns)
|
||||||
|
}
|
||||||
|
res, _ = h.handleCallFunction(ctx, callReq(map[string]any{"name": "admin_only"}))
|
||||||
|
if errCode(t, res) != CodeInvalidArgument {
|
||||||
|
t.Error("an unauthorized function must look unknown")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProcedureFunctionCallShape(t *testing.T) {
|
||||||
|
h, mock, ctx := newTxHarness(t)
|
||||||
|
if err := h.RegisterFunction(Function{
|
||||||
|
Name: "recalc", Procedure: "app.recalc_totals",
|
||||||
|
Params: []FunctionParam{{Name: "account", Type: ParamInteger, Required: true}, {Name: "opts", Type: ParamObject}},
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
mock.ExpectBegin()
|
||||||
|
mock.ExpectQuery(`SELECT \* FROM app\.recalc_totals\(\$1, \$2::jsonb\)`).WithArgs(float64(3), nil).
|
||||||
|
WillReturnRows(sqlmock.NewRows([]string{"total"}).AddRow(10))
|
||||||
|
mock.ExpectCommit()
|
||||||
|
res, _ := h.handleCallFunction(ctx, callReq(map[string]any{"name": "recalc", "arguments": map[string]any{"account": float64(3)}}))
|
||||||
|
if res.IsError {
|
||||||
|
t.Fatal(payload(t, res))
|
||||||
|
}
|
||||||
|
if err := mock.ExpectationsWereMet(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
+2
-350
@@ -1,7 +1,6 @@
|
|||||||
package resolvemcp
|
package resolvemcp
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
"reflect"
|
"reflect"
|
||||||
@@ -10,34 +9,9 @@ import (
|
|||||||
"github.com/mark3labs/mcp-go/mcp"
|
"github.com/mark3labs/mcp-go/mcp"
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
||||||
)
|
)
|
||||||
|
|
||||||
// toolName builds the MCP tool name for a given operation and model.
|
|
||||||
func toolName(operation, schema, entity string) string {
|
|
||||||
if schema == "" {
|
|
||||||
return fmt.Sprintf("%s_%s", operation, entity)
|
|
||||||
}
|
|
||||||
return fmt.Sprintf("%s_%s_%s", operation, schema, entity)
|
|
||||||
}
|
|
||||||
|
|
||||||
// registerModelTools registers the four CRUD tools and resource for a model.
|
|
||||||
func registerModelTools(h *Handler, schema, entity string, model interface{}) {
|
|
||||||
info := buildModelInfo(schema, entity, model)
|
|
||||||
registerReadTool(h, schema, entity, info)
|
|
||||||
registerCreateTool(h, schema, entity, info)
|
|
||||||
registerUpdateTool(h, schema, entity, info)
|
|
||||||
registerDeleteTool(h, schema, entity, info)
|
|
||||||
registerModelResource(h, schema, entity, info)
|
|
||||||
|
|
||||||
logger.Info("[resolvemcp] Registered MCP tools for %s", info.fullName)
|
|
||||||
}
|
|
||||||
|
|
||||||
// --------------------------------------------------------------------------
|
|
||||||
// Model introspection
|
|
||||||
// --------------------------------------------------------------------------
|
|
||||||
|
|
||||||
// modelInfo holds pre-computed metadata for a model used in tool descriptions.
|
// modelInfo holds pre-computed metadata for a model used in tool descriptions.
|
||||||
type modelInfo struct {
|
type modelInfo struct {
|
||||||
fullName string // e.g. "public.users"
|
fullName string // e.g. "public.users"
|
||||||
@@ -247,330 +221,8 @@ func writableColumnNames(cols []columnInfo) []string {
|
|||||||
return names
|
return names
|
||||||
}
|
}
|
||||||
|
|
||||||
// --------------------------------------------------------------------------
|
// parseRequestOptions reads the paging, filter, sort, column and preload arguments shared
|
||||||
// Read tool
|
// by the read tools.
|
||||||
// --------------------------------------------------------------------------
|
|
||||||
|
|
||||||
func registerReadTool(h *Handler, schema, entity string, info modelInfo) {
|
|
||||||
name := toolName("read", schema, entity)
|
|
||||||
|
|
||||||
var descParts []string
|
|
||||||
descParts = append(descParts, fmt.Sprintf("Read records from the '%s' database table.", info.fullName))
|
|
||||||
if info.pkName != "" {
|
|
||||||
descParts = append(descParts, fmt.Sprintf("Primary key: '%s'. Pass it via 'id' to fetch a single record.", info.pkName))
|
|
||||||
}
|
|
||||||
if info.schemaDoc != "" {
|
|
||||||
descParts = append(descParts, info.schemaDoc)
|
|
||||||
}
|
|
||||||
descParts = append(descParts,
|
|
||||||
"Pagination: use 'limit'/'offset' for offset-based paging, or 'cursor_forward'/'cursor_backward' (pass the primary key value of the last/first record on the current page) for cursor-based paging.",
|
|
||||||
"Filtering: each filter object requires 'column' (JSON field name) and 'operator'. Supported operators: = != > < >= <= like ilike in is_null is_not_null. Combine with 'logic_operator': AND (default) or OR.",
|
|
||||||
"Sorting: each sort object requires 'column' and 'direction' (asc or desc).",
|
|
||||||
)
|
|
||||||
if len(info.relationNames) > 0 {
|
|
||||||
descParts = append(descParts, fmt.Sprintf("Preloadable relations: %s. Pass relation name in 'preloads'.", strings.Join(info.relationNames, ", ")))
|
|
||||||
}
|
|
||||||
|
|
||||||
description := strings.Join(descParts, "\n\n")
|
|
||||||
|
|
||||||
filterDesc := `Array of filter objects. Example: [{"column":"status","operator":"=","value":"active"},{"column":"age","operator":">","value":18,"logic_operator":"AND"}]`
|
|
||||||
if len(info.columns) > 0 {
|
|
||||||
filterDesc += fmt.Sprintf(" Available columns: %s.", columnNameList(info.columns))
|
|
||||||
}
|
|
||||||
|
|
||||||
sortDesc := `Array of sort objects. Example: [{"column":"created_at","direction":"desc"}]`
|
|
||||||
if len(info.columns) > 0 {
|
|
||||||
sortDesc += fmt.Sprintf(" Available columns: %s.", columnNameList(info.columns))
|
|
||||||
}
|
|
||||||
|
|
||||||
tool := mcp.NewTool(name,
|
|
||||||
mcp.WithDescription(description),
|
|
||||||
mcp.WithString("id",
|
|
||||||
mcp.Description(fmt.Sprintf("Primary key (%s) of a single record to fetch. Omit to return multiple records.", info.pkName)),
|
|
||||||
),
|
|
||||||
mcp.WithNumber("limit",
|
|
||||||
mcp.Description("Maximum number of records to return per page. Recommended: 10–100."),
|
|
||||||
),
|
|
||||||
mcp.WithNumber("offset",
|
|
||||||
mcp.Description("Number of records to skip (for offset-based pagination). Use with 'limit'."),
|
|
||||||
),
|
|
||||||
mcp.WithString("cursor_forward",
|
|
||||||
mcp.Description(fmt.Sprintf("Cursor for the next page: pass the '%s' value of the last record on the current page. Requires 'sort' to be set.", info.pkName)),
|
|
||||||
),
|
|
||||||
mcp.WithString("cursor_backward",
|
|
||||||
mcp.Description(fmt.Sprintf("Cursor for the previous page: pass the '%s' value of the first record on the current page. Requires 'sort' to be set.", info.pkName)),
|
|
||||||
),
|
|
||||||
mcp.WithArray("columns",
|
|
||||||
mcp.Description(fmt.Sprintf("Columns to include in the result. Omit to return all columns. Available: %s.", columnNameList(info.columns))),
|
|
||||||
),
|
|
||||||
mcp.WithArray("omit_columns",
|
|
||||||
mcp.Description(fmt.Sprintf("Columns to exclude from the result. Available: %s.", columnNameList(info.columns))),
|
|
||||||
),
|
|
||||||
mcp.WithArray("filters",
|
|
||||||
mcp.Description(filterDesc),
|
|
||||||
),
|
|
||||||
mcp.WithArray("sort",
|
|
||||||
mcp.Description(sortDesc),
|
|
||||||
),
|
|
||||||
mcp.WithArray("preloads",
|
|
||||||
mcp.Description(buildPreloadDesc(info)),
|
|
||||||
),
|
|
||||||
)
|
|
||||||
|
|
||||||
h.mcpServer.AddTool(tool, func(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
|
||||||
args := req.GetArguments()
|
|
||||||
id, _ := args["id"].(string)
|
|
||||||
options := parseRequestOptions(args)
|
|
||||||
|
|
||||||
data, metadata, err := h.executeRead(ctx, schema, entity, id, options)
|
|
||||||
if err != nil {
|
|
||||||
return toolError("tool", err), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
return marshalResult(map[string]interface{}{
|
|
||||||
"success": true,
|
|
||||||
"data": data,
|
|
||||||
"metadata": metadata,
|
|
||||||
})
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func buildPreloadDesc(info modelInfo) string {
|
|
||||||
if len(info.relationNames) == 0 {
|
|
||||||
return `Array of relation preload objects. Each object: {"relation":"RelationName"}. No relations defined on this model.`
|
|
||||||
}
|
|
||||||
return fmt.Sprintf(
|
|
||||||
`Array of relation preload objects. Each object: {"relation":"RelationName","columns":["col1","col2"]}. Available relations: %s.`,
|
|
||||||
strings.Join(info.relationNames, ", "),
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
// --------------------------------------------------------------------------
|
|
||||||
// Create tool
|
|
||||||
// --------------------------------------------------------------------------
|
|
||||||
|
|
||||||
func registerCreateTool(h *Handler, schema, entity string, info modelInfo) {
|
|
||||||
name := toolName("create", schema, entity)
|
|
||||||
|
|
||||||
writable := writableColumnNames(info.columns)
|
|
||||||
|
|
||||||
var descParts []string
|
|
||||||
descParts = append(descParts, fmt.Sprintf("Create one or more new records in the '%s' table.", info.fullName))
|
|
||||||
if len(writable) > 0 {
|
|
||||||
descParts = append(descParts, fmt.Sprintf("Writable fields: %s.", strings.Join(writable, ", ")))
|
|
||||||
}
|
|
||||||
if info.pkName != "" {
|
|
||||||
descParts = append(descParts, fmt.Sprintf("The primary key ('%s') is typically auto-generated — omit it unless you need to supply it explicitly.", info.pkName))
|
|
||||||
}
|
|
||||||
descParts = append(descParts,
|
|
||||||
"Pass a single JSON object to 'data' to create one record. Pass an array of objects to create multiple records in a single transaction (all succeed or all fail).",
|
|
||||||
)
|
|
||||||
if info.schemaDoc != "" {
|
|
||||||
descParts = append(descParts, info.schemaDoc)
|
|
||||||
}
|
|
||||||
|
|
||||||
description := strings.Join(descParts, "\n\n")
|
|
||||||
|
|
||||||
dataDesc := "Record fields to create."
|
|
||||||
if len(writable) > 0 {
|
|
||||||
dataDesc += fmt.Sprintf(" Writable fields: %s.", strings.Join(writable, ", "))
|
|
||||||
}
|
|
||||||
dataDesc += " Pass a single object or an array of objects."
|
|
||||||
|
|
||||||
tool := mcp.NewTool(name,
|
|
||||||
mcp.WithDescription(description),
|
|
||||||
mcp.WithObject("data",
|
|
||||||
mcp.Description(dataDesc),
|
|
||||||
mcp.Required(),
|
|
||||||
),
|
|
||||||
)
|
|
||||||
|
|
||||||
h.mcpServer.AddTool(tool, func(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
|
||||||
args := req.GetArguments()
|
|
||||||
data, ok := args["data"]
|
|
||||||
if !ok {
|
|
||||||
return mcp.NewToolResultError("missing required argument: data"), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
result, err := h.executeCreate(ctx, schema, entity, data)
|
|
||||||
if err != nil {
|
|
||||||
return toolError("tool", err), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
return marshalResult(map[string]interface{}{
|
|
||||||
"success": true,
|
|
||||||
"data": result,
|
|
||||||
})
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// --------------------------------------------------------------------------
|
|
||||||
// Update tool
|
|
||||||
// --------------------------------------------------------------------------
|
|
||||||
|
|
||||||
func registerUpdateTool(h *Handler, schema, entity string, info modelInfo) {
|
|
||||||
name := toolName("update", schema, entity)
|
|
||||||
|
|
||||||
writable := writableColumnNames(info.columns)
|
|
||||||
|
|
||||||
var descParts []string
|
|
||||||
descParts = append(descParts, fmt.Sprintf("Update an existing record in the '%s' table.", info.fullName))
|
|
||||||
if info.pkName != "" {
|
|
||||||
descParts = append(descParts, fmt.Sprintf("Identify the record by its primary key ('%s') via the 'id' argument or by including '%s' inside 'data'.", info.pkName, info.pkName))
|
|
||||||
}
|
|
||||||
if len(writable) > 0 {
|
|
||||||
descParts = append(descParts, fmt.Sprintf("Updatable fields: %s.", strings.Join(writable, ", ")))
|
|
||||||
}
|
|
||||||
descParts = append(descParts,
|
|
||||||
"Only non-null, non-empty fields in 'data' are applied — existing values are preserved for fields you omit. Returns the merged record as stored.",
|
|
||||||
)
|
|
||||||
if info.schemaDoc != "" {
|
|
||||||
descParts = append(descParts, info.schemaDoc)
|
|
||||||
}
|
|
||||||
|
|
||||||
description := strings.Join(descParts, "\n\n")
|
|
||||||
|
|
||||||
idDesc := fmt.Sprintf("Primary key ('%s') of the record to update. Can also be included inside 'data'.", info.pkName)
|
|
||||||
|
|
||||||
dataDesc := "Fields to update (non-null, non-empty values are merged into the existing record)."
|
|
||||||
if len(writable) > 0 {
|
|
||||||
dataDesc += fmt.Sprintf(" Updatable fields: %s.", strings.Join(writable, ", "))
|
|
||||||
}
|
|
||||||
|
|
||||||
tool := mcp.NewTool(name,
|
|
||||||
mcp.WithDescription(description),
|
|
||||||
mcp.WithString("id",
|
|
||||||
mcp.Description(idDesc),
|
|
||||||
),
|
|
||||||
mcp.WithObject("data",
|
|
||||||
mcp.Description(dataDesc),
|
|
||||||
mcp.Required(),
|
|
||||||
),
|
|
||||||
)
|
|
||||||
|
|
||||||
h.mcpServer.AddTool(tool, func(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
|
||||||
args := req.GetArguments()
|
|
||||||
id, _ := args["id"].(string)
|
|
||||||
|
|
||||||
data, ok := args["data"]
|
|
||||||
if !ok {
|
|
||||||
return mcp.NewToolResultError("missing required argument: data"), nil
|
|
||||||
}
|
|
||||||
dataMap, ok := data.(map[string]interface{})
|
|
||||||
if !ok {
|
|
||||||
return mcp.NewToolResultError("data must be an object"), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
result, err := h.executeUpdate(ctx, schema, entity, id, dataMap)
|
|
||||||
if err != nil {
|
|
||||||
return toolError("tool", err), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
return marshalResult(map[string]interface{}{
|
|
||||||
"success": true,
|
|
||||||
"data": result,
|
|
||||||
})
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// --------------------------------------------------------------------------
|
|
||||||
// Delete tool
|
|
||||||
// --------------------------------------------------------------------------
|
|
||||||
|
|
||||||
func registerDeleteTool(h *Handler, schema, entity string, info modelInfo) {
|
|
||||||
name := toolName("delete", schema, entity)
|
|
||||||
|
|
||||||
descParts := []string{
|
|
||||||
fmt.Sprintf("Delete a record from the '%s' table by its primary key.", info.fullName),
|
|
||||||
}
|
|
||||||
if info.pkName != "" {
|
|
||||||
descParts = append(descParts, fmt.Sprintf("Pass the '%s' value of the record to delete via the 'id' argument.", info.pkName))
|
|
||||||
}
|
|
||||||
descParts = append(descParts, "Returns the deleted record. This operation is irreversible.")
|
|
||||||
|
|
||||||
description := strings.Join(descParts, " ")
|
|
||||||
|
|
||||||
tool := mcp.NewTool(name,
|
|
||||||
mcp.WithDescription(description),
|
|
||||||
mcp.WithString("id",
|
|
||||||
mcp.Description(fmt.Sprintf("Primary key ('%s') of the record to delete.", info.pkName)),
|
|
||||||
mcp.Required(),
|
|
||||||
),
|
|
||||||
)
|
|
||||||
|
|
||||||
h.mcpServer.AddTool(tool, func(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
|
||||||
args := req.GetArguments()
|
|
||||||
id, _ := args["id"].(string)
|
|
||||||
|
|
||||||
result, err := h.executeDelete(ctx, schema, entity, id)
|
|
||||||
if err != nil {
|
|
||||||
return toolError("tool", err), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
return marshalResult(map[string]interface{}{
|
|
||||||
"success": true,
|
|
||||||
"data": result,
|
|
||||||
})
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// --------------------------------------------------------------------------
|
|
||||||
// Resource registration
|
|
||||||
// --------------------------------------------------------------------------
|
|
||||||
|
|
||||||
func registerModelResource(h *Handler, schema, entity string, info modelInfo) {
|
|
||||||
resourceURI := info.fullName
|
|
||||||
|
|
||||||
var resourceDesc strings.Builder
|
|
||||||
fmt.Fprintf(&resourceDesc, "Database table: %s", info.fullName)
|
|
||||||
if info.pkName != "" {
|
|
||||||
fmt.Fprintf(&resourceDesc, " (primary key: %s)", info.pkName)
|
|
||||||
}
|
|
||||||
if info.schemaDoc != "" {
|
|
||||||
resourceDesc.WriteString("\n\n")
|
|
||||||
resourceDesc.WriteString(info.schemaDoc)
|
|
||||||
}
|
|
||||||
|
|
||||||
resource := mcp.NewResource(
|
|
||||||
resourceURI,
|
|
||||||
entity,
|
|
||||||
mcp.WithResourceDescription(resourceDesc.String()),
|
|
||||||
mcp.WithMIMEType("application/json"),
|
|
||||||
)
|
|
||||||
|
|
||||||
h.mcpServer.AddResource(resource, func(ctx context.Context, req mcp.ReadResourceRequest) ([]mcp.ResourceContents, error) {
|
|
||||||
limit := 100
|
|
||||||
options := common.RequestOptions{Limit: &limit}
|
|
||||||
|
|
||||||
data, metadata, err := h.executeRead(ctx, schema, entity, "", options)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
payload := map[string]interface{}{
|
|
||||||
"data": data,
|
|
||||||
"metadata": metadata,
|
|
||||||
}
|
|
||||||
jsonBytes, err := json.Marshal(payload)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("error marshaling resource: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return []mcp.ResourceContents{
|
|
||||||
mcp.TextResourceContents{
|
|
||||||
URI: req.Params.URI,
|
|
||||||
MIMEType: "application/json",
|
|
||||||
Text: string(jsonBytes),
|
|
||||||
},
|
|
||||||
}, nil
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// --------------------------------------------------------------------------
|
|
||||||
// Argument parsing helpers
|
|
||||||
// --------------------------------------------------------------------------
|
|
||||||
|
|
||||||
// parseRequestOptions converts raw MCP tool arguments into common.RequestOptions.
|
|
||||||
func parseRequestOptions(args map[string]interface{}) common.RequestOptions {
|
func parseRequestOptions(args map[string]interface{}) common.RequestOptions {
|
||||||
options := common.RequestOptions{}
|
options := common.RequestOptions{}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,259 @@
|
|||||||
|
package resolvemcp
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"reflect"
|
||||||
|
"regexp"
|
||||||
|
"sort"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security"
|
||||||
|
)
|
||||||
|
|
||||||
|
// previewRows is how many matched primary keys a preview lists.
|
||||||
|
const previewRows = 10
|
||||||
|
|
||||||
|
var filterColumnRe = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*$`)
|
||||||
|
|
||||||
|
// writeOperators are the filter operators a filter write accepts.
|
||||||
|
var writeOperators = map[string]bool{
|
||||||
|
"eq": true, "=": true, "neq": true, "!=": true, "<>": true, "gt": true, ">": true, "gte": true, ">=": true,
|
||||||
|
"lt": true, "<": true, "lte": true, "<=": true, "like": true, "ilike": true, "in": true,
|
||||||
|
"is_null": true, "is_not_null": true,
|
||||||
|
}
|
||||||
|
|
||||||
|
// validateWriteFilters requires at least one filter and that every filter is usable. Reads
|
||||||
|
// silently drop filters they cannot apply; a write must not, because a dropped filter widens
|
||||||
|
// the set of rows the write touches.
|
||||||
|
func validateWriteFilters(model interface{}, filters []common.FilterOption) error {
|
||||||
|
if len(filters) == 0 {
|
||||||
|
return invalidArg("provide an id or at least one filter")
|
||||||
|
}
|
||||||
|
v := common.NewColumnValidator(model)
|
||||||
|
for _, f := range filters {
|
||||||
|
if !filterColumnRe.MatchString(f.Column) || !v.IsValidColumn(f.Column) {
|
||||||
|
return invalidArg("unknown filter column %q", truncate(f.Column))
|
||||||
|
}
|
||||||
|
if !writeOperators[strings.ToLower(f.Operator)] {
|
||||||
|
return invalidArg("unsupported filter operator %q", truncate(f.Operator))
|
||||||
|
}
|
||||||
|
op := strings.ToLower(f.Operator)
|
||||||
|
if op != "is_null" && op != "is_not_null" && f.Value == nil {
|
||||||
|
return invalidArg("filter on %q needs a value", f.Column)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// whereRequest describes a filter-based update or delete.
|
||||||
|
type whereRequest struct {
|
||||||
|
schema, entity string
|
||||||
|
op string // "update" or "delete"
|
||||||
|
filters []common.FilterOption
|
||||||
|
data map[string]interface{} // update only
|
||||||
|
dryRun bool
|
||||||
|
confirmToken string
|
||||||
|
}
|
||||||
|
|
||||||
|
// whereResult is what a filter write returns to the client.
|
||||||
|
type whereResult struct {
|
||||||
|
DryRun bool `json:"dry_run,omitempty"`
|
||||||
|
RequiresConfirm bool `json:"requires_confirmation,omitempty"`
|
||||||
|
Matched int `json:"matched"`
|
||||||
|
Preview []interface{} `json:"preview,omitempty"`
|
||||||
|
ConfirmToken string `json:"confirm_token,omitempty"`
|
||||||
|
ExpiresInSec int `json:"expires_in_seconds,omitempty"`
|
||||||
|
Affected int `json:"affected,omitempty"`
|
||||||
|
IDs []interface{} `json:"ids,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// executeWhere runs a filter-based write behind the guardrails: filters are validated, the
|
||||||
|
// matching rows are counted inside the transaction (aborting above Config.MaxWriteRows), and
|
||||||
|
// without dry_run the write only happens with a confirm token from a preview of the same
|
||||||
|
// request that matched the same rows.
|
||||||
|
func (h *Handler) executeWhere(ctx context.Context, req whereRequest) (_ *whereResult, retErr error) {
|
||||||
|
defer recoverPanic(&retErr)
|
||||||
|
ctx, cancel := h.callContext(ctx)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
model, err := h.registry.GetModelByEntity(req.schema, req.entity)
|
||||||
|
if err != nil {
|
||||||
|
return nil, invalidArg("model not found: %s", buildModelName(req.schema, req.entity))
|
||||||
|
}
|
||||||
|
unwrapped, err := common.ValidateAndUnwrapModel(model)
|
||||||
|
if err != nil {
|
||||||
|
return nil, errInternal
|
||||||
|
}
|
||||||
|
model = unwrapped.Model
|
||||||
|
if err := validateWriteFilters(model, req.filters); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
pkName := reflection.GetPrimaryKeyName(model)
|
||||||
|
if pkName == "" {
|
||||||
|
return nil, invalidArg("table has no primary key; filter writes are not available")
|
||||||
|
}
|
||||||
|
tableName := h.getTableName(req.schema, req.entity, model)
|
||||||
|
ctx = withRequestData(h.withModelRules(ctx, req.schema, req.entity), req.schema, req.entity, tableName, model, unwrapped.ModelPtr)
|
||||||
|
|
||||||
|
var setCols map[string]interface{}
|
||||||
|
if req.op == "update" {
|
||||||
|
if setCols, err = writeColumns(model, req.data); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
for col := range setCols {
|
||||||
|
if strings.EqualFold(col, pkName) {
|
||||||
|
delete(setCols, col)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(setCols) == 0 {
|
||||||
|
return nil, invalidArg("no updatable fields in data")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
hookCtx := &HookContext{
|
||||||
|
Context: ctx, Handler: h, Schema: req.schema, Entity: req.entity, Model: model,
|
||||||
|
Operation: req.op, Tx: h.db,
|
||||||
|
}
|
||||||
|
if req.op == "update" {
|
||||||
|
hookCtx.Data = req.data
|
||||||
|
}
|
||||||
|
if err := h.hooks.Execute(BeforeHandle, hookCtx); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
user := callerKey(ctx)
|
||||||
|
table := buildModelName(req.schema, req.entity)
|
||||||
|
res := &whereResult{}
|
||||||
|
err = h.runInTx(ctx, hookCtx, func(tx common.Database) error {
|
||||||
|
before, after := BeforeUpdate, AfterUpdate
|
||||||
|
if req.op == "delete" {
|
||||||
|
before, after = BeforeDelete, AfterDelete
|
||||||
|
}
|
||||||
|
if err := h.hooks.Execute(before, hookCtx); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if req.op == "update" {
|
||||||
|
// A hook (column security) may have narrowed the payload.
|
||||||
|
if m, ok := hookCtx.Data.(map[string]interface{}); ok {
|
||||||
|
if setCols, err = writeColumns(model, m); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
for col := range setCols {
|
||||||
|
if strings.EqualFold(col, pkName) {
|
||||||
|
delete(setCols, col)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(setCols) == 0 {
|
||||||
|
return invalidArg("no updatable fields in data")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
ids, err := h.matchRows(ctx, tx, hookCtx, model, unwrapped.ModelType, tableName, pkName, req.filters)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
res.Matched = len(ids)
|
||||||
|
for i := 0; i < len(ids) && i < previewRows; i++ {
|
||||||
|
res.Preview = append(res.Preview, ids[i])
|
||||||
|
}
|
||||||
|
|
||||||
|
if req.dryRun {
|
||||||
|
res.DryRun = true
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
binding, err := bindingHash(req.op, req.filters, req.data, ids)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if req.confirmToken == "" {
|
||||||
|
tok, err := h.confirms.issue(user, table, req.op, binding, h.config.ConfirmTTL)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
res.RequiresConfirm = true
|
||||||
|
res.ConfirmToken = tok
|
||||||
|
res.ExpiresInSec = int(h.config.ConfirmTTL / time.Second)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if err := h.confirms.consume(req.confirmToken, user, table, req.op, binding); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if len(ids) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
inList := make([]string, len(ids))
|
||||||
|
for i := range ids {
|
||||||
|
inList[i] = "?"
|
||||||
|
}
|
||||||
|
cond := fmt.Sprintf("%s IN (%s)", common.QuoteIdent(pkName), strings.Join(inList, ", "))
|
||||||
|
var affected int64
|
||||||
|
if req.op == "update" {
|
||||||
|
r, err := tx.NewUpdate().Table(tableName).SetMap(setCols).Where(cond, ids...).Exec(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("error updating records: %w", err)
|
||||||
|
}
|
||||||
|
affected = r.RowsAffected()
|
||||||
|
} else {
|
||||||
|
r, err := tx.NewDelete().Table(tableName).Where(cond, ids...).Exec(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("delete error: %w", err)
|
||||||
|
}
|
||||||
|
affected = r.RowsAffected()
|
||||||
|
}
|
||||||
|
if int(affected) != len(ids) {
|
||||||
|
// Rows changed under us: roll back rather than report a partial write.
|
||||||
|
return invalidArg("matched rows changed during the write; repeat the preview")
|
||||||
|
}
|
||||||
|
res.Affected = int(affected)
|
||||||
|
res.IDs = ids
|
||||||
|
hookCtx.Result = map[string]interface{}{"ids": ids, "count": len(ids)}
|
||||||
|
return h.hooks.Execute(after, hookCtx)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return res, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// matchRows returns the primary keys of the rows the filters select, narrowed by the BeforeScan
|
||||||
|
// hooks (row security). It fails when more than Config.MaxWriteRows match.
|
||||||
|
func (h *Handler) matchRows(ctx context.Context, tx common.Database, hookCtx *HookContext, model interface{}, modelType reflect.Type, tableName, pkName string, filters []common.FilterOption) ([]interface{}, error) {
|
||||||
|
sliceType := reflect.SliceOf(reflect.PointerTo(modelType))
|
||||||
|
dest := reflect.New(sliceType)
|
||||||
|
q := tx.NewSelect().Model(dest.Interface())
|
||||||
|
if provider, ok := reflect.New(modelType).Interface().(common.TableNameProvider); !ok || provider.TableName() == "" {
|
||||||
|
q = q.Table(tableName)
|
||||||
|
}
|
||||||
|
q = h.applyFilters(q.Column(pkName), filters, model)
|
||||||
|
hookCtx.Query = q
|
||||||
|
if err := h.hooks.Execute(BeforeScan, hookCtx); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if err := hookCtx.Query.Limit(h.config.MaxWriteRows + 1).ScanModel(ctx); err != nil {
|
||||||
|
return nil, fmt.Errorf("error matching records: %w", err)
|
||||||
|
}
|
||||||
|
rows := dest.Elem()
|
||||||
|
if rows.Len() > h.config.MaxWriteRows {
|
||||||
|
return nil, NewClientError(CodeLimitExceeded, fmt.Sprintf("the filters match more than %d rows; narrow them", h.config.MaxWriteRows))
|
||||||
|
}
|
||||||
|
ids := make([]interface{}, 0, rows.Len())
|
||||||
|
for i := 0; i < rows.Len(); i++ {
|
||||||
|
ids = append(ids, reflection.GetPrimaryKeyValue(rows.Index(i).Interface()))
|
||||||
|
}
|
||||||
|
sort.Slice(ids, func(i, j int) bool { return fmt.Sprint(ids[i]) < fmt.Sprint(ids[j]) })
|
||||||
|
return ids, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// callerKey identifies the caller for confirm-token binding.
|
||||||
|
func callerKey(ctx context.Context) string {
|
||||||
|
if uc, ok := security.GetUserContext(ctx); ok && uc != nil {
|
||||||
|
return fmt.Sprintf("%d/%s", uc.UserID, uc.UserName)
|
||||||
|
}
|
||||||
|
return "anonymous"
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user