diff --git a/pkg/common/adapters/database/pgsql.go b/pkg/common/adapters/database/pgsql.go index 9616601..ef5a3fe 100644 --- a/pkg/common/adapters/database/pgsql.go +++ b/pkg/common/adapters/database/pgsql.go @@ -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 } diff --git a/pkg/common/adapters/database/pgsql_test.go b/pkg/common/adapters/database/pgsql_test.go index e5b5192..dac8cb3 100644 --- a/pkg/common/adapters/database/pgsql_test.go +++ b/pkg/common/adapters/database/pgsql_test.go @@ -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()) +} diff --git a/pkg/common/tx_guard_test.go b/pkg/common/tx_guard_test.go index a8977cb..192c790 100644 --- a/pkg/common/tx_guard_test.go +++ b/pkg/common/tx_guard_test.go @@ -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() diff --git a/pkg/resolvemcp/confirm.go b/pkg/resolvemcp/confirm.go new file mode 100644 index 0000000..c898079 --- /dev/null +++ b/pkg/resolvemcp/confirm.go @@ -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 +} diff --git a/pkg/resolvemcp/functions.go b/pkg/resolvemcp/functions.go new file mode 100644 index 0000000..32bb57d --- /dev/null +++ b/pkg/resolvemcp/functions.go @@ -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 +} diff --git a/pkg/resolvemcp/handler.go b/pkg/resolvemcp/handler.go index 584eebc..ce8ae88 100644 --- a/pkg/resolvemcp/handler.go +++ b/pkg/resolvemcp/handler.go @@ -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 } diff --git a/pkg/resolvemcp/hooks.go b/pkg/resolvemcp/hooks.go index 7d59d25..78c0854 100644 --- a/pkg/resolvemcp/hooks.go +++ b/pkg/resolvemcp/hooks.go @@ -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 diff --git a/pkg/resolvemcp/meta.go b/pkg/resolvemcp/meta.go new file mode 100644 index 0000000..6c862d4 --- /dev/null +++ b/pkg/resolvemcp/meta.go @@ -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}) +} diff --git a/pkg/resolvemcp/meta_test.go b/pkg/resolvemcp/meta_test.go new file mode 100644 index 0000000..157858d --- /dev/null +++ b/pkg/resolvemcp/meta_test.go @@ -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) + } +} diff --git a/pkg/resolvemcp/tools.go b/pkg/resolvemcp/tools.go index f8901bf..d198fb7 100644 --- a/pkg/resolvemcp/tools.go +++ b/pkg/resolvemcp/tools.go @@ -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{} diff --git a/pkg/resolvemcp/writewhere.go b/pkg/resolvemcp/writewhere.go new file mode 100644 index 0000000..826bd3a --- /dev/null +++ b/pkg/resolvemcp/writewhere.go @@ -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" +}