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:
Hein
2026-10-01 13:40:00 +02:00
parent 276c3814d8
commit e49c3a916e
11 changed files with 1562 additions and 370 deletions
+13 -14
View File
@@ -5,7 +5,9 @@ import (
"database/sql"
"fmt"
"reflect"
"regexp"
"sort"
"strconv"
"strings"
"sync"
"time"
@@ -767,6 +769,9 @@ func (p *PgSQLInsertQuery) Scan(ctx context.Context, dest interface{}) (err erro
return nil
}
// placeholderRe matches a numbered SQL parameter such as $12.
var placeholderRe = regexp.MustCompile(`\$\d+`)
// PgSQLUpdateQuery implements UpdateQuery for PostgreSQL
type PgSQLUpdateQuery struct {
db *sql.DB
@@ -897,23 +902,17 @@ func (p *PgSQLUpdateQuery) Exec(ctx context.Context) (res common.Result, err err
p.tableName,
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 {
shift := len(setArgs)
updatedWhereClauses := make([]string, 0, len(p.whereClauses))
for _, whereClause := range p.whereClauses {
// Find and replace parameter placeholders
updatedClause := whereClause
paramNum := i
// 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
updatedWhereClauses = append(updatedWhereClauses, placeholderRe.ReplaceAllStringFunc(whereClause, func(m string) string {
n, _ := strconv.Atoi(m[1:])
return fmt.Sprintf("$%d", n+shift)
}))
}
p.whereClauses = updatedWhereClauses
}
@@ -627,3 +627,23 @@ func TestRawSQL(t *testing.T) {
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())
}
+4 -3
View File
@@ -27,12 +27,13 @@ var allowedPoolHookTx = map[string]int{
"resolvespec/handler.go": 1,
"websocketspec/handler.go": 1,
"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.
var allowedPoolQuery = map[string]int{
"resolvemcp/annotation.go": 2, // tool annotations, not a data request
}
var allowedPoolQuery = map[string]int{}
func guardedFiles(t *testing.T) map[string][]string {
t.Helper()
+83
View File
@@ -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
}
+280
View File
@@ -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
}
+6 -3
View File
@@ -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
}
+6
View File
@@ -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
+438
View File
@@ -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})
}
+451
View File
@@ -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
View File
@@ -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{}
+259
View File
@@ -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"
}