mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-02 19:41:57 +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:
@@ -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
|
||||
oauth2Regs []oauth2Registration
|
||||
oauthSrv *security.OAuthServer
|
||||
functions functionRegistry
|
||||
confirms *confirmStore
|
||||
}
|
||||
|
||||
// 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(),
|
||||
mcpServer: server.NewMCPServer("resolvemcp", "1.0.0"),
|
||||
config: cfg.withDefaults(),
|
||||
confirms: newConfirmStore(),
|
||||
name: "resolvemcp",
|
||||
version: "1.0.0",
|
||||
}
|
||||
registerMetaTools(h)
|
||||
if cfg.EnableAnnotations {
|
||||
registerAnnotationTool(h)
|
||||
}
|
||||
@@ -165,13 +169,13 @@ func requestBaseURL(r *http.Request) string {
|
||||
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 {
|
||||
fullName := buildModelName(schema, entity)
|
||||
if err := h.registry.RegisterModel(fullName, model); err != nil {
|
||||
return err
|
||||
}
|
||||
registerModelTools(h, schema, entity, model)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -187,7 +191,6 @@ func (h *Handler) RegisterModelWithRules(schema, entity string, model interface{
|
||||
if err := reg.RegisterModelWithRules(fullName, model, rules); err != nil {
|
||||
return err
|
||||
}
|
||||
registerModelTools(h, schema, entity, model)
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
@@ -34,6 +34,12 @@ const (
|
||||
BeforeDelete HookType = "before_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
|
||||
// (including the second short transaction for post-commit work). hookCtx.Tx is
|
||||
// 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
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"reflect"
|
||||
@@ -10,34 +9,9 @@ import (
|
||||
"github.com/mark3labs/mcp-go/mcp"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
"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.
|
||||
type modelInfo struct {
|
||||
fullName string // e.g. "public.users"
|
||||
@@ -247,330 +221,8 @@ func writableColumnNames(cols []columnInfo) []string {
|
||||
return names
|
||||
}
|
||||
|
||||
// --------------------------------------------------------------------------
|
||||
// Read tool
|
||||
// --------------------------------------------------------------------------
|
||||
|
||||
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.
|
||||
// parseRequestOptions reads the paging, filter, sort, column and preload arguments shared
|
||||
// by the read tools.
|
||||
func parseRequestOptions(args map[string]interface{}) 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