Compare commits

..
4 Commits
Author SHA1 Message Date
warkanum e4c4315f4b fix(filters): qualify model columns to avoid ambiguity with joins
Qualify plain model columns with the main table alias in resolvespec
filters so preload joins no longer cause SQLSTATE 42702, and accept
<table>.<column> from clients in resolvespec and restheadspec instead
of silently dropping the filter. Add join/preload tests for
resolvespec, restheadspec and funcspec.
2026-10-07 20:58:44 +02:00
warkanum 3efd539e0f style: gofmt test files 2026-10-07 20:58:44 +02:00
Hein 5a3a1df3c8 feat(resolvemcp)!: make the server read-only by default
Tests / Integration Tests (push) Skipped
Build , Vet Test, and Lint / Build (push) Successful in 1m33s
Tests / Unit Tests (push) Successful in 2m1s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 2m36s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 2m38s
Build , Vet Test, and Lint / Lint Code (push) Successful in 2m49s
Tests / Race Detector (push) Successful in 4m31s
BREAKING CHANGE: Config.ReadOnly is now a *bool and unset means read-only.
Use ReadOnly: resolvemcp.Bool(false) to enable insert/update/delete,
annotations and function calls. Adds the Bool helper and updates docs.
2026-10-07 14:17:56 +02:00
Hein 4ed9506ad2 feat(resolvemcp): add read-only mode and function allowlist
- Config.ReadOnly disables insert/update/delete/annotation tools, reports only
  select in list_tables/describe_table and tells the agent it cannot write
- Config.AllowFunctionCalls keeps function tools on a read-only server
- Config.AllowedFunctions limits list_functions/call_function to named
  functions (empty allows all); others are reported as unknown
- reflect read-only mode in the usage guide and exported catalogue
2026-10-07 14:15:14 +02:00
23 changed files with 1069 additions and 93 deletions
@@ -1,3 +1,4 @@
//go:build integration
// +build integration
package database
@@ -18,11 +19,11 @@ import (
// Integration test models
type IntegrationUser struct {
ID int `db:"id"`
Name string `db:"name"`
Email string `db:"email"`
Age int `db:"age"`
CreatedAt time.Time `db:"created_at"`
ID int `db:"id"`
Name string `db:"name"`
Email string `db:"email"`
Age int `db:"age"`
CreatedAt time.Time `db:"created_at"`
Posts []*IntegrationPost `bun:"rel:has-many,join:id=user_id"`
}
@@ -46,10 +47,10 @@ func (p IntegrationPost) TableName() string {
}
type IntegrationComment struct {
ID int `db:"id"`
Content string `db:"content"`
PostID int `db:"post_id"`
CreatedAt time.Time `db:"created_at"`
ID int `db:"id"`
Content string `db:"content"`
PostID int `db:"post_id"`
CreatedAt time.Time `db:"created_at"`
Post *IntegrationPost `bun:"rel:belongs-to,join:post_id=id"`
}
+5 -5
View File
@@ -26,11 +26,11 @@ func (u TestUser) TableName() string {
}
type TestPost struct {
ID int `db:"id"`
Title string `db:"title"`
Content string `db:"content"`
UserID int `db:"user_id"`
User *TestUser `bun:"rel:belongs-to,join:user_id=id"`
ID int `db:"id"`
Title string `db:"title"`
Content string `db:"content"`
UserID int `db:"user_id"`
User *TestUser `bun:"rel:belongs-to,join:user_id=id"`
Comments []TestComment `bun:"rel:has-many,join:id=post_id"`
}
+77
View File
@@ -0,0 +1,77 @@
package common
import (
"regexp"
"strings"
"github.com/bitechdev/ResolveSpec/pkg/reflection"
)
var rePlainIdent = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*$`)
// MainTableAlias returns the alias the main table gets in the SELECT:
// the model's TableAlias() when provided, otherwise the bare table name.
func MainTableAlias(model interface{}, tableName string) string {
if p, ok := model.(TableAliasProvider); ok {
if a := p.TableAlias(); a != "" {
return a
}
}
return reflection.ExtractTableNameOnly(tableName)
}
func isModelSQLColumn(model interface{}, column string) bool {
for _, c := range reflection.GetSQLModelColumns(model) {
if strings.EqualFold(c, column) {
return true
}
}
return false
}
// QualifyModelColumn returns "alias"."column" when column is a plain identifier
// that exists on the model, so it stays unambiguous once joins are added.
// Anything else (expressions, JSON paths, other tables' columns) is returned unchanged.
func QualifyModelColumn(model interface{}, alias, column string) string {
if alias == "" || model == nil || !rePlainIdent.MatchString(column) || !isModelSQLColumn(model, column) {
return column
}
return QuoteIdent(alias) + "." + QuoteIdent(column)
}
// StripMainTablePrefix turns "<prefix>.<column>" into "<column>" when prefix is
// one of the main table's names/aliases and column exists on the model.
// Any other input is returned unchanged.
func StripMainTablePrefix(model interface{}, column string, prefixes ...string) string {
idx := strings.Index(column, ".")
if idx <= 0 || model == nil {
return column
}
prefix := strings.Trim(column[:idx], `"`)
col := strings.Trim(column[idx+1:], `"`)
if !rePlainIdent.MatchString(prefix) || !rePlainIdent.MatchString(col) {
return column
}
for _, p := range prefixes {
if p != "" && strings.EqualFold(p, prefix) && isModelSQLColumn(model, col) {
return col
}
}
return column
}
// StripMainTablePrefixFromFilters applies StripMainTablePrefix to every filter column.
func StripMainTablePrefixFromFilters(model interface{}, filters []FilterOption, prefixes ...string) {
for i := range filters {
filters[i].Column = StripMainTablePrefix(model, filters[i].Column, prefixes...)
}
}
// NormalizeMainTableFilters rewrites "<main table or alias>.<column>" filter
// columns on opts to the bare model column (on a copy, never the caller's
// slice) so the column validator keeps them.
func NormalizeMainTableFilters(model interface{}, tableName string, opts *RequestOptions) {
opts.Filters = append([]FilterOption(nil), opts.Filters...)
StripMainTablePrefixFromFilters(model, opts.Filters,
MainTableAlias(model, tableName), reflection.ExtractTableNameOnly(tableName))
}
+55
View File
@@ -0,0 +1,55 @@
package common
import "testing"
type qualifyModel struct {
ID int64 `bun:"id,pk"`
Name string `bun:"name"`
}
func TestQualifyModelColumn(t *testing.T) {
m := qualifyModel{}
cases := []struct{ alias, col, want string }{
{"t", "name", `"t"."name"`},
{"t", "NAME", `"t"."NAME"`},
{"t", "missing", "missing"},
{"t", "data->>'x'", "data->>'x'"},
{"t", "rel.name", "rel.name"},
{"", "name", "name"},
}
for _, c := range cases {
if got := QualifyModelColumn(m, c.alias, c.col); got != c.want {
t.Errorf("Qualify(%q,%q) = %q, want %q", c.alias, c.col, got, c.want)
}
}
if got := QualifyModelColumn(nil, "t", "name"); got != "name" {
t.Errorf("nil model: %q", got)
}
}
func TestStripMainTablePrefix(t *testing.T) {
m := qualifyModel{}
cases := []struct{ col, want string }{
{"province_state.name", "name"},
{`"province_state"."name"`, "name"},
{"PROVINCE_STATE.name", "name"},
{"rel_rid_country.name", "rel_rid_country.name"},
{"province_state.missing", "province_state.missing"},
{"name", "name"},
}
for _, c := range cases {
if got := StripMainTablePrefix(m, c.col, "province_state"); got != c.want {
t.Errorf("Strip(%q) = %q, want %q", c.col, got, c.want)
}
}
}
func TestStripMainTablePrefix_ValidatorKeepsFilter(t *testing.T) {
m := qualifyModel{}
filters := []FilterOption{{Column: "province_state.name", Operator: "eq", Value: "x"}}
StripMainTablePrefixFromFilters(m, filters, "province_state")
out := NewColumnValidator(m).FilterRequestOptions(RequestOptions{Filters: filters})
if len(out.Filters) != 1 || out.Filters[0].Column != "name" {
t.Fatalf("filters = %+v", out.Filters)
}
}
+43 -35
View File
@@ -25,10 +25,10 @@ func newMockDatabase() *mockDatabase {
}
}
func (m *mockDatabase) NewSelect() SelectQuery { return &mockSelectQuery{} }
func (m *mockDatabase) NewInsert() InsertQuery { return &mockInsertQuery{db: m} }
func (m *mockDatabase) NewUpdate() UpdateQuery { return &mockUpdateQuery{db: m} }
func (m *mockDatabase) NewDelete() DeleteQuery { return &mockDeleteQuery{db: m} }
func (m *mockDatabase) NewSelect() SelectQuery { return &mockSelectQuery{} }
func (m *mockDatabase) NewInsert() InsertQuery { return &mockInsertQuery{db: m} }
func (m *mockDatabase) NewUpdate() UpdateQuery { return &mockUpdateQuery{db: m} }
func (m *mockDatabase) NewDelete() DeleteQuery { return &mockDeleteQuery{db: m} }
func (m *mockDatabase) RunInTransaction(ctx context.Context, fn func(Database) error) error {
return fn(m)
}
@@ -57,27 +57,31 @@ func (m *mockDatabase) DriverName() string {
// Mock SelectQuery
type mockSelectQuery struct{}
func (m *mockSelectQuery) Model(model interface{}) SelectQuery { return m }
func (m *mockSelectQuery) Table(name string) SelectQuery { return m }
func (m *mockSelectQuery) Column(columns ...string) SelectQuery { return m }
func (m *mockSelectQuery) ColumnExpr(query string, args ...interface{}) SelectQuery { return m }
func (m *mockSelectQuery) Where(condition string, args ...interface{}) SelectQuery { return m }
func (m *mockSelectQuery) WhereOr(query string, args ...interface{}) SelectQuery { return m }
func (m *mockSelectQuery) Join(query string, args ...interface{}) SelectQuery { return m }
func (m *mockSelectQuery) LeftJoin(query string, args ...interface{}) SelectQuery { return m }
func (m *mockSelectQuery) Model(model interface{}) SelectQuery { return m }
func (m *mockSelectQuery) Table(name string) SelectQuery { return m }
func (m *mockSelectQuery) Column(columns ...string) SelectQuery { return m }
func (m *mockSelectQuery) ColumnExpr(query string, args ...interface{}) SelectQuery { return m }
func (m *mockSelectQuery) Where(condition string, args ...interface{}) SelectQuery { return m }
func (m *mockSelectQuery) WhereOr(query string, args ...interface{}) SelectQuery { return m }
func (m *mockSelectQuery) Join(query string, args ...interface{}) SelectQuery { return m }
func (m *mockSelectQuery) LeftJoin(query string, args ...interface{}) SelectQuery { return m }
func (m *mockSelectQuery) Preload(relation string, conditions ...interface{}) SelectQuery { return m }
func (m *mockSelectQuery) PreloadRelation(relation string, apply ...func(SelectQuery) SelectQuery) SelectQuery { return m }
func (m *mockSelectQuery) JoinRelation(relation string, apply ...func(SelectQuery) SelectQuery) SelectQuery { return m }
func (m *mockSelectQuery) Order(order string) SelectQuery { return m }
func (m *mockSelectQuery) OrderExpr(order string, args ...interface{}) SelectQuery { return m }
func (m *mockSelectQuery) Limit(n int) SelectQuery { return m }
func (m *mockSelectQuery) Offset(n int) SelectQuery { return m }
func (m *mockSelectQuery) Group(group string) SelectQuery { return m }
func (m *mockSelectQuery) PreloadRelation(relation string, apply ...func(SelectQuery) SelectQuery) SelectQuery {
return m
}
func (m *mockSelectQuery) JoinRelation(relation string, apply ...func(SelectQuery) SelectQuery) SelectQuery {
return m
}
func (m *mockSelectQuery) Order(order string) SelectQuery { return m }
func (m *mockSelectQuery) OrderExpr(order string, args ...interface{}) SelectQuery { return m }
func (m *mockSelectQuery) Limit(n int) SelectQuery { return m }
func (m *mockSelectQuery) Offset(n int) SelectQuery { return m }
func (m *mockSelectQuery) Group(group string) SelectQuery { return m }
func (m *mockSelectQuery) Having(condition string, args ...interface{}) SelectQuery { return m }
func (m *mockSelectQuery) Scan(ctx context.Context, dest interface{}) error { return nil }
func (m *mockSelectQuery) ScanModel(ctx context.Context) error { return nil }
func (m *mockSelectQuery) Count(ctx context.Context) (int, error) { return 0, nil }
func (m *mockSelectQuery) Exists(ctx context.Context) (bool, error) { return false, nil }
func (m *mockSelectQuery) Scan(ctx context.Context, dest interface{}) error { return nil }
func (m *mockSelectQuery) ScanModel(ctx context.Context) error { return nil }
func (m *mockSelectQuery) Count(ctx context.Context) (int, error) { return 0, nil }
func (m *mockSelectQuery) Exists(ctx context.Context) (bool, error) { return false, nil }
// Mock InsertQuery
type mockInsertQuery struct {
@@ -98,9 +102,9 @@ func (m *mockInsertQuery) Value(column string, value interface{}) InsertQuery {
m.values[column] = value
return m
}
func (m *mockInsertQuery) OnConflict(action string) InsertQuery { return m }
func (m *mockInsertQuery) OnConflict(action string) InsertQuery { return m }
func (m *mockInsertQuery) ExcludeColumn(columns ...string) InsertQuery { return m }
func (m *mockInsertQuery) Returning(columns ...string) InsertQuery { return m }
func (m *mockInsertQuery) Returning(columns ...string) InsertQuery { return m }
func (m *mockInsertQuery) Exec(ctx context.Context) (Result, error) {
m.db.insertCalls = append(m.db.insertCalls, m.values)
m.db.lastID++
@@ -132,8 +136,8 @@ func (m *mockUpdateQuery) SetMap(values map[string]interface{}) UpdateQuery {
return m
}
func (m *mockUpdateQuery) Where(condition string, args ...interface{}) UpdateQuery { return m }
func (m *mockUpdateQuery) ExcludeColumn(columns ...string) UpdateQuery { return m }
func (m *mockUpdateQuery) Returning(columns ...string) UpdateQuery { return m }
func (m *mockUpdateQuery) ExcludeColumn(columns ...string) UpdateQuery { return m }
func (m *mockUpdateQuery) Returning(columns ...string) UpdateQuery { return m }
func (m *mockUpdateQuery) Exec(ctx context.Context) (Result, error) {
// Record the update call
m.db.updateCalls = append(m.db.updateCalls, m.setValues)
@@ -171,9 +175,13 @@ func (m *mockResult) RowsAffected() int64 { return m.rowsAffected }
type mockModelRegistry struct{}
func (m *mockModelRegistry) GetModel(name string) (interface{}, error) { return nil, nil }
func (m *mockModelRegistry) GetModelByEntity(schema, entity string) (interface{}, error) { return nil, nil }
func (m *mockModelRegistry) GetModelByEntity(schema, entity string) (interface{}, error) {
return nil, nil
}
func (m *mockModelRegistry) RegisterModel(name string, model interface{}) error { return nil }
func (m *mockModelRegistry) GetAllModels() map[string]interface{} { return make(map[string]interface{}) }
func (m *mockModelRegistry) GetAllModels() map[string]interface{} {
return make(map[string]interface{})
}
// Mock RelationshipInfoProvider
type mockRelationshipProvider struct {
@@ -198,9 +206,9 @@ func (m *mockRelationshipProvider) RegisterRelation(modelTypeName, relationName
// Test Models
type Department struct {
ID int64 `json:"id" bun:"id,pk"`
Name string `json:"name"`
Employees []*Employee `json:"employees,omitempty"`
ID int64 `json:"id" bun:"id,pk"`
Name string `json:"name"`
Employees []*Employee `json:"employees,omitempty"`
}
func (d Department) TableName() string { return "departments" }
@@ -227,9 +235,9 @@ func (t Task) TableName() string { return "tasks" }
func (t Task) GetIDName() string { return "ID" }
type Comment struct {
ID int64 `json:"id" bun:"id,pk"`
Text string `json:"text"`
TaskID int64 `json:"task_id"`
ID int64 `json:"id" bun:"id,pk"`
Text string `json:"text"`
TaskID int64 `json:"task_id"`
}
func (c Comment) TableName() string { return "comments" }
+92
View File
@@ -0,0 +1,92 @@
package funcspec
import (
"database/sql"
"strings"
"testing"
"github.com/uptrace/bun/driver/sqliteshim"
)
const joinedBaseSQL = "SELECT p.id, p.name, c.name AS country_name FROM province_state p LEFT JOIN country c ON c.id = p.rid_country"
// funcspec filters are appended to author-written SQL with no model, so columns
// are used verbatim: clients disambiguate by sending "alias.column", and the
// dot must survive ValidSQL and every filter path.
func TestApplyFilters_QualifiedColumnsPreservedWithJoin(t *testing.T) {
h := NewHandler(&MockDatabase{})
cases := []struct {
name string
params *RequestParameters
want string
}{
{"field filter", &RequestParameters{FieldFilters: map[string]string{"p.name": "abc"}}, "p.name = abc"},
{"search filter", &RequestParameters{SearchFilters: map[string]string{"p.name": "abc"}}, "CAST(p.name AS TEXT) ILIKE '%abc%'"},
{"search op eq", &RequestParameters{SearchOps: map[string]FilterOperator{"p.name": {Operator: "eq", Value: "abc", Logic: "AND"}}}, "p.name = 'abc'"},
{"search op contains", &RequestParameters{SearchOps: map[string]FilterOperator{"p.name": {Operator: "contains", Value: "abc", Logic: "AND"}}}, "CAST(p.name AS TEXT) ILIKE '%abc%'"},
}
for _, c := range cases {
t.Run(c.name, func(t *testing.T) {
got := h.ApplyFilters(joinedBaseSQL, c.params)
if !strings.Contains(got, c.want) {
t.Fatalf("SQL %q does not contain %q", got, c.want)
}
if !strings.Contains(got, "LEFT JOIN country") || !strings.Contains(got, " WHERE ") {
t.Fatalf("join/where lost: %q", got)
}
})
}
}
func TestApplyFilters_WithAndWithoutJoin_RealQuery(t *testing.T) {
sqldb, err := sql.Open(sqliteshim.ShimName, "file:funcspecjoin?mode=memory&cache=private")
if err != nil {
t.Fatal(err)
}
defer sqldb.Close()
for _, stmt := range []string{
"CREATE TABLE country (id INTEGER PRIMARY KEY, name TEXT)",
"CREATE TABLE province_state (id INTEGER PRIMARY KEY, name TEXT, rid_country INTEGER)",
"INSERT INTO country VALUES (1,'Abcland'),(2,'Other')",
"INSERT INTO province_state VALUES (1,'abc one',1),(2,'xyz two',1),(3,'nope',2)",
} {
if _, err := sqldb.Exec(stmt); err != nil {
t.Fatal(err)
}
}
h := NewHandler(&MockDatabase{})
count := func(base string, params *RequestParameters) (int, error) {
rows, err := sqldb.Query(h.ApplyFilters(base, params))
if err != nil {
return 0, err
}
defer rows.Close()
n := 0
for rows.Next() {
n++
}
return n, rows.Err()
}
const noJoin = "SELECT p.id, p.name FROM province_state p"
params := func(col string) *RequestParameters {
return &RequestParameters{SearchOps: map[string]FilterOperator{col: {Operator: "eq", Value: "nope", Logic: "AND"}}}
}
if n, err := count(noJoin, params("name")); err != nil || n != 1 {
t.Fatalf("no join, unqualified: n=%d err=%v", n, err)
}
if n, err := count(joinedBaseSQL, params("p.name")); err != nil || n != 1 {
t.Fatalf("join, qualified: n=%d err=%v", n, err)
}
if n, err := count(joinedBaseSQL, params("c.name")); err != nil || n != 0 {
t.Fatalf("join, joined-table column: n=%d err=%v", n, err)
}
// Documented limit: with a join in the author's SQL, an unqualified shared
// column is ambiguous and the client must send "alias.column".
if _, err := count(joinedBaseSQL, params("name")); err == nil || !strings.Contains(strings.ToLower(err.Error()), "ambiguous") {
t.Fatalf("expected ambiguous error for unqualified shared column, got %v", err)
}
}
+32
View File
@@ -16,6 +16,8 @@ import (
handler := resolvemcp.NewHandlerWithGORM(db, resolvemcp.Config{
BaseURL: "http://localhost:8080",
BasePath: "/mcp",
// Read-only by default; uncomment to allow writes:
// ReadOnly: resolvemcp.Bool(false),
})
securityList, _ := security.NewSecurityList(provider)
@@ -442,6 +444,35 @@ The text appears in `list_tables` and `describe_table`. The server also sends a
`handler.ExportCatalog(path)` writes the usage guide, tools, limits and every table (columns, types, keys, relations, allowed operations, descriptions) to disk, JSON for a `.json` path and Markdown otherwise. The file is replaced atomically. It lists every table with at least one allowed operation, regardless of caller, so keep it out of public directories. Call it after registering models (for example at startup, or from a `go generate` step).
## Read-only mode
The server is **read-only unless you enable writes**: `Config.ReadOnly` is a `*bool` and an unset (nil) value means on. To allow inserts, updates and deletes:
```go
handler := resolvemcp.NewHandlerWithGORM(db, resolvemcp.Config{ReadOnly: resolvemcp.Bool(false)})
```
While read-only is on:
- The insert, update, delete and annotation tools are not registered, so the agent never sees them. `list_functions`/`call_function` are off too, because a registered function may change data, unless you set `AllowFunctionCalls` (below).
- `list_tables` and `describe_table` report only `select`; `describe_table` also sets `read_only: true` and lists no writable columns.
- The MCP server instructions (and the exported catalogue) say the server is read-only and tell the agent not to attempt writes.
- A write that reaches a handler anyway is refused with a `forbidden` error ("this server is read-only: writes are disabled").
### Function calls and the allowlist
```go
// Read-only server that may still run two named functions
resolvemcp.Config{
// ReadOnly is on by default
AllowFunctionCalls: true, // keep list_functions / call_function on a read-only server
AllowedFunctions: []string{"report_totals", "search_customers"},
}
```
- `AllowFunctionCalls` only matters while read-only is on; with writes enabled (`ReadOnly: resolvemcp.Bool(false)`), functions are always available. Set it only for functions that do not change data.
- `AllowedFunctions` works in either mode. When empty, every registered function is allowed. When set, only the named functions are listed and callable; any other is reported as `unknown function`, so its existence is not revealed. Per-function `Authorize` still applies on top.
## MCP Tools
Fixed set, independent of the models. `table` is `schema.entity`. Errors return `{"success":false,"error":{"code","message"}}` with codes `invalid_argument`, `not_found`, `forbidden`, `limit_exceeded`, `internal` (internal details are logged, the client gets a reference id).
@@ -702,6 +733,7 @@ The handler resolves table names in priority order:
## Breaking changes
- The server is read-only by default. Writes (insert/update/delete), annotations and function calls need `Config{ReadOnly: resolvemcp.Bool(false)}` (function calls can also be kept on a read-only server with `AllowFunctionCalls`).
- Per-model tools (`read_/create_/update_/delete_{schema}_{entity}`) and per-model resources are gone; use the meta tools.
- `Setup*` / `NewSSEServer` / `NewStreamableHTTPHandler` take a `*security.SecurityList` and require authentication. `OptionalAuth*` helpers were removed; `*Unauthenticated` variants exist for explicit opt-out.
- `resolvespec_annotate` is opt-in via `Config.EnableAnnotations`.
+30 -3
View File
@@ -24,12 +24,35 @@ const usageGuide = `This server exposes database tables through a fixed set of t
5. Use list_functions / call_function for registered functions.
Read the error message when a call fails: it says which argument was wrong.`
// readOnlyGuide replaces usageGuide on a read-only server.
const readOnlyGuide = `This server exposes database tables through a fixed set of tools. It is READ-ONLY: you cannot insert, update or delete data or write annotations, and no tool for that exists. Do not attempt a write; tell the user it is not possible through this server.
1. Call list_tables to see the tables you may read and what they hold.
2. Call describe_table for a table before using it: columns, types, primary key, relations (preloadable) and limits.
3. Read with select_table (filters, sort, columns, preloads). Results are paged; use limit/offset or cursors, and include_count only when you need a total.
Read the error message when a call fails: it says which argument was wrong.`
// readOnlyFunctionsGuide is the extra step of a read-only server that still allows functions.
const readOnlyFunctionsGuide = `
4. Use list_functions / call_function for the registered functions. Only call functions that fit a read-only server; the server decides what is allowed.`
// guideFor returns the usage guide for the server mode.
func guideFor(readOnly, functions bool) string {
if !readOnly {
return usageGuide
}
if functions {
return readOnlyGuide + readOnlyFunctionsGuide
}
return readOnlyGuide
}
// Catalog is a snapshot of what the server offers: the usage guide, the tools, the limits
// and every table with its columns, relations, allowed operations and descriptions.
type Catalog struct {
GeneratedAt time.Time `json:"generated_at"`
Server string `json:"server"`
Version string `json:"version"`
ReadOnly bool `json:"read_only"`
Guide string `json:"guide"`
Limits CatalogLimits `json:"limits"`
Tools []CatalogTool `json:"tools"`
@@ -122,7 +145,8 @@ func (h *Handler) BuildCatalog() Catalog {
GeneratedAt: time.Now().UTC(),
Server: h.name,
Version: h.version,
Guide: usageGuide,
ReadOnly: h.config.readOnly,
Guide: guideFor(h.config.readOnly, h.config.AllowFunctionCalls),
Limits: CatalogLimits{
DefaultLimit: h.config.DefaultLimit,
MaxLimit: h.config.MaxLimit,
@@ -143,7 +167,7 @@ func (h *Handler) BuildCatalog() Catalog {
for name, model := range h.registry.GetAllModels() {
schema, entity, _ := splitTable(name)
rules := h.modelRules(schema, entity)
ops := opsFor(rules)
ops := h.opsFor(rules)
if len(ops) == 0 {
continue
}
@@ -155,7 +179,7 @@ func (h *Handler) BuildCatalog() Catalog {
for mt != nil && (mt.Kind() == reflect.Pointer || mt.Kind() == reflect.Slice) {
mt = mt.Elem()
}
if mt != nil && mt.Kind() == reflect.Struct {
if !h.config.readOnly && mt != nil && mt.Kind() == reflect.Struct {
for k := range reflectionJSONColumns(mt) {
writable[k] = true
}
@@ -234,6 +258,9 @@ func (h *Handler) ExportCatalog(path string) error {
func (c Catalog) Markdown() string {
var sb strings.Builder
fmt.Fprintf(&sb, "# %s API catalogue\n\nGenerated %s.\n\n", c.Server, c.GeneratedAt.Format(time.RFC3339))
if c.ReadOnly {
sb.WriteString("**This server is read-only.**\n\n")
}
sb.WriteString("## How to use\n\n" + c.Guide + "\n\n")
fmt.Fprintf(&sb, "## Limits\n\ndefault limit %d, max limit %d, max offset %d, max batch %d, max preload depth %d, max rows per filter write %d.\n\n",
c.Limits.DefaultLimit, c.Limits.MaxLimit, c.Limits.MaxOffset, c.Limits.MaxBatch, c.Limits.MaxPreloadDepth, c.Limits.MaxWriteRows)
+3
View File
@@ -17,6 +17,9 @@
//
// The same guide is sent to MCP clients as the server instructions.
//
// The server is read-only by default (Config.ReadOnly nil means on); set
// ReadOnly: resolvemcp.Bool(false) to enable the write tools.
//
// # Setting it up
//
// handler := resolvemcp.NewHandlerWithGORM(db, resolvemcp.Config{BaseURL: "http://localhost:8080"})
+13 -1
View File
@@ -108,6 +108,15 @@ func (h *Handler) function(name string) (Function, bool) {
return f, ok
}
// functionAllowed reports whether Config.AllowedFunctions lets the function through.
func (h *Handler) functionAllowed(name string) bool {
if h.allowedFns == nil {
return true
}
_, ok := h.allowedFns[name]
return ok
}
// visibleFunctions returns the functions the caller may call, sorted by name.
func (h *Handler) visibleFunctions(ctx context.Context) []Function {
h.functions.mu.RLock()
@@ -119,6 +128,9 @@ func (h *Handler) visibleFunctions(ctx context.Context) []Function {
sort.Slice(out, func(i, j int) bool { return out[i].Name < out[j].Name })
visible := out[:0]
for _, f := range out {
if !h.functionAllowed(f.Name) {
continue
}
if f.Authorize == nil || f.Authorize(ctx) == nil {
visible = append(visible, f)
}
@@ -236,7 +248,7 @@ func (h *Handler) executeCall(ctx context.Context, name string, rawArgs map[stri
defer cancel()
f, ok := h.function(name)
if !ok {
if !ok || !h.functionAllowed(name) {
return nil, invalidArg("unknown function %q", truncate(name))
}
hookCtx := &HookContext{Context: ctx, Handler: h, Entity: name, Operation: "call_function", Tx: h.db}
+9 -2
View File
@@ -24,6 +24,7 @@ import (
// Handler exposes registered database models as MCP tools and resources.
type Handler struct {
allowedFns map[string]struct{} // nil: every function is allowed
db common.Database
registry common.ModelRegistry
hooks *HookRegistry
@@ -43,14 +44,20 @@ func NewHandler(db common.Database, registry common.ModelRegistry, cfg Config) *
db: db,
registry: registry,
hooks: NewHookRegistry(),
mcpServer: server.NewMCPServer("resolvemcp", "1.0.0", server.WithInstructions(usageGuide)),
mcpServer: server.NewMCPServer("resolvemcp", "1.0.0", server.WithInstructions(guideFor(cfg.withDefaults().readOnly, cfg.AllowFunctionCalls))),
config: cfg.withDefaults(),
confirms: newConfirmStore(),
name: "resolvemcp",
version: "1.0.0",
}
if len(cfg.AllowedFunctions) > 0 {
h.allowedFns = make(map[string]struct{}, len(cfg.AllowedFunctions))
for _, n := range cfg.AllowedFunctions {
h.allowedFns[n] = struct{}{}
}
}
registerMetaTools(h)
if cfg.EnableAnnotations {
if cfg.EnableAnnotations && !h.config.readOnly {
registerAnnotationTool(h)
}
return h
+27 -6
View File
@@ -56,6 +56,16 @@ func registerMetaTools(h *Handler) {
mcp.WithBoolean("include_count", mcp.Description("Also return the total number of matching rows (slower on large tables).")),
), h.handleSelect)
if !h.config.readOnly {
registerWriteTools(h, tableArg, idArg, filtersArg, dryRunArg, confirmArg)
}
if !h.config.readOnly || h.config.AllowFunctionCalls {
registerFunctionTools(h, readOnly)
}
}
// registerWriteTools adds the tools that change table rows.
func registerWriteTools(h *Handler, tableArg, idArg, filtersArg, dryRunArg, confirmArg mcp.ToolOption) {
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.")),
@@ -73,7 +83,10 @@ func registerMetaTools(h *Handler) {
mcp.WithDestructiveHintAnnotation(true),
tableArg, idArg, filtersArg, dryRunArg, confirmArg,
), h.handleDelete)
}
// registerFunctionTools adds list_functions and call_function.
func registerFunctionTools(h *Handler, readOnly mcp.ToolOption) {
h.mcpServer.AddTool(mcp.NewTool("list_functions", readOnly,
mcp.WithDescription("List the functions you can call with call_function, with their parameters.")),
h.handleListFunctions)
@@ -111,11 +124,15 @@ func (h *Handler) modelRules(schema, entity string) modelregistry.ModelRules {
return modelregistry.DefaultModelRules()
}
func opsFor(r modelregistry.ModelRules) []string {
// opsFor lists the operations the rules allow. A read-only server allows select only.
func (h *Handler) opsFor(r modelregistry.ModelRules) []string {
var ops []string
if r.CanRead {
ops = append(ops, opSelect)
}
if h.config.readOnly {
return ops
}
if r.CanCreate {
ops = append(ops, opInsert)
}
@@ -140,9 +157,12 @@ func (h *Handler) resolveTable(args map[string]any, op string) (schema, entity s
if _, err := h.registry.GetModelByEntity(schema, entity); err != nil {
return "", "", invalidArg("unknown table %q; see list_tables", truncate(table))
}
if op != "" && op != opSelect && h.config.readOnly {
return "", "", NewClientError(CodeForbidden, "this server is read-only: writes are disabled")
}
if op != "" {
allowed := false
for _, o := range opsFor(h.modelRules(schema, entity)) {
for _, o := range h.opsFor(h.modelRules(schema, entity)) {
if o == op {
allowed = true
}
@@ -163,7 +183,7 @@ func (h *Handler) handleListTables(ctx context.Context, _ mcp.CallToolRequest) (
var tables []table
for name := range h.registry.GetAllModels() {
schema, entity, _ := splitTable(name)
if ops := opsFor(h.modelRules(schema, entity)); len(ops) > 0 {
if ops := h.opsFor(h.modelRules(schema, entity)); len(ops) > 0 {
tables = append(tables, table{Table: name, Description: h.modelDocs(schema, entity).Description, Operations: ops})
}
}
@@ -181,7 +201,7 @@ func (h *Handler) handleDescribeTable(_ context.Context, req mcp.CallToolRequest
return toolError("describe_table", invalidArg("unknown table")), nil
}
rules := h.modelRules(schema, entity)
if len(opsFor(rules)) == 0 {
if len(h.opsFor(rules)) == 0 {
return toolError("describe_table", invalidArg("unknown table %q; see list_tables", buildModelName(schema, entity))), nil
}
info := buildModelInfo(schema, entity, model)
@@ -192,7 +212,7 @@ func (h *Handler) handleDescribeTable(_ context.Context, req mcp.CallToolRequest
modelType = modelType.Elem()
}
writable := map[string]bool{}
if modelType != nil && modelType.Kind() == reflect.Struct {
if !h.config.readOnly && modelType != nil && modelType.Kind() == reflect.Struct {
for jsonKey := range reflection.BuildJSONToDBColumnMap(modelType) {
writable[jsonKey] = true
}
@@ -230,7 +250,8 @@ func (h *Handler) handleDescribeTable(_ context.Context, req mcp.CallToolRequest
"columns": cols,
"relations": info.relationNames,
"writable_columns": writableNames,
"operations": opsFor(rules),
"operations": h.opsFor(rules),
"read_only": h.config.readOnly,
"filter_operators": filterOperators,
"limits": map[string]any{
"default_limit": h.config.DefaultLimit,
+168
View File
@@ -0,0 +1,168 @@
package resolvemcp
import (
"context"
"strings"
"testing"
"github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
)
func newReadOnlyHandler(t *testing.T) *Handler {
t.Helper()
h := NewHandler(database.NewPgSQLAdapter(nil), modelregistry.NewModelRegistry(),
Config{EnableAnnotations: true})
if err := h.RegisterModel("public", "items", &docItem{}); err != nil {
t.Fatal(err)
}
return h
}
func TestReadOnlyToolSet(t *testing.T) {
h := newReadOnlyHandler(t)
tools := h.mcpServer.ListTools()
for _, name := range []string{"list_tables", "describe_table", "select_table"} {
if tools[name] == nil {
t.Errorf("read tool %s missing", name)
}
}
for _, name := range []string{"insert_into_table", "update_table", "delete_from_table", "call_function", "list_functions", annotationToolName} {
if tools[name] != nil {
t.Errorf("tool %s must not be registered on a read-only server", name)
}
}
}
func TestReadOnlyRefusesWritesAndReportsIt(t *testing.T) {
h := newReadOnlyHandler(t)
ctx := context.Background()
args := map[string]any{"table": "public.items", "data": map[string]any{"name": "x"}, "id": 1}
for name, fn := range map[string]func() map[string]any{
"insert": func() map[string]any { r, _ := h.handleInsert(ctx, callReq(args)); return payload(t, r) },
"update": func() map[string]any { r, _ := h.handleUpdate(ctx, callReq(args)); return payload(t, r) },
"delete": func() map[string]any { r, _ := h.handleDelete(ctx, callReq(args)); return payload(t, r) },
} {
e, _ := fn()["error"].(map[string]any)
if e["code"] != CodeForbidden || !strings.Contains(e["message"].(string), "read-only") {
t.Errorf("%s: error = %v", name, e)
}
}
res, _ := h.handleListTables(ctx, callReq(nil))
tb := payload(t, res)["tables"].([]any)[0].(map[string]any)
if ops := tb["operations"].([]any); len(ops) != 1 || ops[0] != opSelect {
t.Errorf("list_tables operations = %v", ops)
}
res, _ = h.handleDescribeTable(ctx, callReq(map[string]any{"table": "public.items"}))
p := payload(t, res)
if p["read_only"] != true {
t.Errorf("describe_table read_only = %v", p["read_only"])
}
if w, _ := p["writable_columns"].([]any); len(w) != 0 {
t.Errorf("writable_columns = %v", w)
}
cat := h.BuildCatalog()
if !cat.ReadOnly || !strings.Contains(cat.Guide, "READ-ONLY") || !strings.Contains(cat.Markdown(), "read-only") {
t.Error("catalogue must say the server is read-only")
}
for _, c := range cat.Tables[0].Columns {
if c.Writable {
t.Errorf("column %s marked writable", c.Name)
}
}
if !strings.Contains(guideFor(true, false), "READ-ONLY") || strings.Contains(guideFor(false, false), "READ-ONLY") {
t.Error("guideFor")
}
}
func newFnHandler(t *testing.T, cfg Config) *Handler {
t.Helper()
h := NewHandler(database.NewPgSQLAdapter(nil), modelregistry.NewModelRegistry(), cfg)
for _, name := range []string{"alpha", "beta"} {
name := name
err := h.RegisterFunction(Function{Name: name, Handler: func(context.Context, common.Database, map[string]any) (any, error) {
return name, nil
}})
if err != nil {
t.Fatal(err)
}
}
return h
}
func TestReadOnlyAllowFunctionCalls(t *testing.T) {
h := newFnHandler(t, Config{AllowFunctionCalls: true})
tools := h.mcpServer.ListTools()
if tools["list_functions"] == nil || tools["call_function"] == nil {
t.Error("function tools must be registered")
}
if tools["insert_into_table"] != nil || tools["update_table"] != nil {
t.Error("write tools must stay off")
}
if g := guideFor(true, true); !strings.Contains(g, "READ-ONLY") || !strings.Contains(g, "call_function") {
t.Error("guide must mention functions")
}
if h.mcpServer.ListTools()["call_function"] == nil {
t.Error("call_function missing")
}
}
func TestAllowedFunctions(t *testing.T) {
ctx := context.Background()
for name, tc := range map[string]struct {
allowed []string
visible []string
}{
"empty allows all": {nil, []string{"alpha", "beta"}},
"only listed": {[]string{"beta"}, []string{"beta"}},
"unknown name": {[]string{"zzz"}, nil},
} {
h := newFnHandler(t, Config{AllowedFunctions: tc.allowed})
var got []string
for _, f := range h.visibleFunctions(ctx) {
got = append(got, f.Name)
}
if strings.Join(got, ",") != strings.Join(tc.visible, ",") {
t.Errorf("%s: visible = %v, want %v", name, got, tc.visible)
}
for _, fn := range []string{"alpha", "beta"} {
listed := false
for _, v := range tc.visible {
listed = listed || v == fn
}
if h.functionAllowed(fn) != listed {
t.Errorf("%s: functionAllowed(%s) = %v, want %v", name, fn, !listed, listed)
}
if !listed {
// refused before any database work, and indistinguishable from a missing function
if _, err := h.executeCall(ctx, fn, nil); err == nil || !strings.Contains(err.Error(), "unknown function") {
t.Errorf("%s: %s must be reported unknown, err=%v", name, fn, err)
}
}
}
}
}
func TestReadOnlyDefaultsOnAndCanBeDisabled(t *testing.T) {
if !(Config{}).withDefaults().readOnly {
t.Error("ReadOnly must default to on")
}
if !(Config{ReadOnly: Bool(true)}).withDefaults().readOnly {
t.Error("explicit true")
}
if (Config{ReadOnly: Bool(false)}).withDefaults().readOnly {
t.Error("Bool(false) must enable writes")
}
h := NewHandler(database.NewPgSQLAdapter(nil), modelregistry.NewModelRegistry(), Config{ReadOnly: Bool(false)})
if h.mcpServer.ListTools()["insert_into_table"] == nil {
t.Error("write tools must register when ReadOnly is Bool(false)")
}
if strings.Contains(h.BuildCatalog().Guide, "READ-ONLY") {
t.Error("guide must not claim read-only")
}
}
+26
View File
@@ -51,6 +51,28 @@ type Config struct {
// host, with at most 32 distinct base URLs cached; prefer setting BaseURL.
AllowedHosts []string
// ReadOnly disables every write and is ON when left nil: set it to Bool(false) to allow
// writes. When on, the insert, update, delete and annotation tools are not registered,
// list_tables and describe_table report only the select operation (no writable columns),
// a write attempted anyway is refused with a "forbidden" error, and the server
// instructions tell the agent it cannot write. list_functions/call_function are also
// off, because a registered function may change data, unless AllowFunctionCalls is set.
ReadOnly *bool
// readOnly is ReadOnly after defaults (nil means true).
readOnly bool
// AllowFunctionCalls keeps list_functions and call_function available on a ReadOnly
// server. Only set it for functions that do not change data; pair it with
// AllowedFunctions to name them. It has no effect when writes are enabled (ReadOnly set to Bool(false)) (functions are
// always available then).
AllowFunctionCalls bool
// AllowedFunctions restricts list_functions and call_function to the named functions.
// Empty allows every registered function. A function outside the list is reported as
// unknown, so its existence is not revealed.
AllowedFunctions []string
// EnableAnnotations registers the resolvespec_annotate tool. Off by default: annotations
// are free text that agents read back, so enabling the tool opens a write channel into
// agent-visible text. When on, every call runs the BeforeHandle hooks (operation
@@ -58,6 +80,9 @@ type Config struct {
EnableAnnotations bool
}
// Bool returns a pointer to v, for the optional boolean fields of Config.
func Bool(v bool) *bool { return &v }
// withDefaults fills the zero limit fields.
func (c Config) withDefaults() Config {
def := func(v *int, d int) {
@@ -74,6 +99,7 @@ func (c Config) withDefaults() Config {
if c.DefaultLimit > c.MaxLimit {
c.DefaultLimit = c.MaxLimit
}
c.readOnly = c.ReadOnly == nil || *c.ReadOnly
if c.QueryTimeout <= 0 {
c.QueryTimeout = 30 * time.Second
}
+1 -1
View File
@@ -191,7 +191,7 @@ func TestAnnotationToolIsOptIn(t *testing.T) {
if h.mcpServer.GetTool(annotationToolName) != nil {
t.Fatal("annotation tool must be off by default")
}
on := NewHandler(h.db, modelregistry.NewModelRegistry(), Config{EnableAnnotations: true})
on := NewHandler(h.db, modelregistry.NewModelRegistry(), Config{EnableAnnotations: true, ReadOnly: Bool(false)})
if on.mcpServer.GetTool(annotationToolName) == nil {
t.Fatal("annotation tool missing when enabled")
}
+1 -1
View File
@@ -29,7 +29,7 @@ func newTxHarness(t *testing.T) (*Handler, sqlmock.Sqlmock, context.Context) {
// connection and fails on the context timeout.
db.SetMaxOpenConns(1)
t.Cleanup(func() { _ = db.Close() })
h := NewHandler(database.NewPgSQLAdapter(db), modelregistry.NewModelRegistry(), Config{})
h := NewHandler(database.NewPgSQLAdapter(db), modelregistry.NewModelRegistry(), Config{ReadOnly: Bool(false)})
if err := h.RegisterModel("public", "items", &txItem{}); err != nil {
t.Fatal(err)
}
+183
View File
@@ -0,0 +1,183 @@
package resolvespec
import (
"context"
"database/sql"
"strings"
"testing"
"github.com/uptrace/bun"
"github.com/uptrace/bun/dialect/sqlitedialect"
"github.com/uptrace/bun/driver/sqliteshim"
"github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
)
type joinCountry struct {
bun.BaseModel `bun:"table:country,alias:country"`
ID int64 `bun:"id,pk"`
Name string `bun:"name"`
}
type joinProvince struct {
bun.BaseModel `bun:"table:province_state,alias:province_state"`
ID int64 `bun:"id,pk"`
Name string `bun:"name"`
Abbreviation string `bun:"abbreviation"`
RidCountry int64 `bun:"rid_country"`
}
// SQLite has no ILIKE; the ilike operator is covered by SQL-string tests, and
// LIKE here proves the CAST(... AS TEXT) wrapping executes with qualified columns.
func setupJoinDB(t *testing.T) *bun.DB {
t.Helper()
sqldb, err := sql.Open(sqliteshim.ShimName, "file:filterjoin?mode=memory&cache=private")
if err != nil {
t.Fatal(err)
}
db := bun.NewDB(sqldb, sqlitedialect.New())
t.Cleanup(func() { _ = db.Close() })
ctx := context.Background()
for _, m := range []interface{}{(*joinCountry)(nil), (*joinProvince)(nil)} {
if _, err := db.NewCreateTable().Model(m).IfNotExists().Exec(ctx); err != nil {
t.Fatal(err)
}
}
if _, err := db.NewInsert().Model(&[]joinCountry{{ID: 1, Name: "Abcland"}, {ID: 2, Name: "Other"}}).Exec(ctx); err != nil {
t.Fatal(err)
}
if _, err := db.NewInsert().Model(&[]joinProvince{
{ID: 1, Name: "abc one", Abbreviation: "A1", RidCountry: 1},
{ID: 2, Name: "xyz two", Abbreviation: "ABC", RidCountry: 1},
{ID: 3, Name: "nope", Abbreviation: "N3", RidCountry: 2},
}).Exec(ctx); err != nil {
t.Fatal(err)
}
return db
}
const joinSQL = "LEFT JOIN country AS rel_rid_country ON rel_rid_country.id = province_state.rid_country"
func countWith(t *testing.T, db *bun.DB, withJoin bool, filters []common.FilterOption, alias string) (int, error) {
t.Helper()
h := &Handler{}
var q common.SelectQuery = database.NewBunAdapter(db).NewSelect().Model(&[]*joinProvince{})
if withJoin {
q = q.Join(joinSQL)
}
q = h.applyFilters(q, filters, &joinProvince{}, alias)
return q.Count(context.Background())
}
func TestFilters_AmbiguousColumnWithJoin_RealQuery(t *testing.T) {
db := setupJoinDB(t)
filters := []common.FilterOption{
{Column: "name", Operator: "like", Value: "%abc%"},
{Column: "abbreviation", Operator: "like", Value: "%abc%", LogicOperator: "OR"},
}
t.Run("unqualified with join is ambiguous (the bug)", func(t *testing.T) {
_, err := countWith(t, db, true, filters, "")
if err == nil || !strings.Contains(strings.ToLower(err.Error()), "ambiguous") {
t.Fatalf("expected ambiguous column error, got %v", err)
}
})
t.Run("qualified with join", func(t *testing.T) {
n, err := countWith(t, db, true, filters, "province_state")
if err != nil {
t.Fatal(err)
}
if n != 2 {
t.Fatalf("count = %d, want 2", n)
}
})
t.Run("qualified without join", func(t *testing.T) {
n, err := countWith(t, db, false, filters, "province_state")
if err != nil {
t.Fatal(err)
}
if n != 2 {
t.Fatalf("count = %d, want 2", n)
}
})
t.Run("unqualified without join still works", func(t *testing.T) {
n, err := countWith(t, db, false, filters, "")
if err != nil {
t.Fatal(err)
}
if n != 2 {
t.Fatalf("count = %d, want 2", n)
}
})
}
func TestFilters_AllOperatorsWithJoin_RealQuery(t *testing.T) {
db := setupJoinDB(t)
cases := []struct {
name string
filter common.FilterOption
want int
}{
{"eq", common.FilterOption{Column: "name", Operator: "eq", Value: "nope"}, 1},
{"neq", common.FilterOption{Column: "name", Operator: "neq", Value: "nope"}, 2},
{"like", common.FilterOption{Column: "name", Operator: "like", Value: "abc%"}, 1},
{"in", common.FilterOption{Column: "name", Operator: "in", Value: []string{"nope", "xyz two"}}, 2},
{"gt", common.FilterOption{Column: "id", Operator: "gt", Value: 1}, 2},
{"gte", common.FilterOption{Column: "id", Operator: "gte", Value: 2}, 2},
{"lt", common.FilterOption{Column: "id", Operator: "lt", Value: 3}, 2},
{"lte", common.FilterOption{Column: "id", Operator: "lte", Value: 1}, 1},
}
for _, c := range cases {
t.Run(c.name, func(t *testing.T) {
n, err := countWith(t, db, true, []common.FilterOption{c.filter}, "province_state")
if err != nil {
t.Fatal(err)
}
if n != c.want {
t.Fatalf("count = %d, want %d", n, c.want)
}
})
}
}
// A filter on the joined table's own column, sent already qualified, must pass through untouched.
func TestFilters_JoinedTableColumnPassesThrough_RealQuery(t *testing.T) {
db := setupJoinDB(t)
n, err := countWith(t, db, true, []common.FilterOption{
{Column: "rel_rid_country.name", Operator: "eq", Value: "Abcland"},
}, "province_state")
if err != nil {
t.Fatal(err)
}
if n != 2 {
t.Fatalf("count = %d, want 2", n)
}
}
// Client sends "province_state.name": it must survive validation and be applied.
func TestFilters_ClientQualifiedColumn_NotDropped(t *testing.T) {
db := setupJoinDB(t)
model := &joinProvince{}
opts := common.RequestOptions{Filters: []common.FilterOption{
{Column: "province_state.name", Operator: "like", Value: "%abc%"},
}}
common.NormalizeMainTableFilters(model, "public.province_state", &opts)
opts = common.NewColumnValidator(model).FilterRequestOptions(opts)
if len(opts.Filters) != 1 {
t.Fatalf("filter was dropped: %+v", opts.Filters)
}
n, err := countWith(t, db, true, opts.Filters, "province_state")
if err != nil {
t.Fatal(err)
}
if n != 1 {
t.Fatalf("count = %d, want 1 (filter must narrow the result)", n)
}
}
+65
View File
@@ -0,0 +1,65 @@
package resolvespec
import (
"testing"
"github.com/bitechdev/ResolveSpec/pkg/common"
)
func TestBuildFilterConditionAlias_QualifiesModelColumns(t *testing.T) {
h := &Handler{}
model := jsonColModel{}
tests := []struct {
name string
filter common.FilterOption
want string
}{
{"eq", common.FilterOption{Column: "name", Operator: "eq", Value: "x"}, `"province_state"."name" = ?`},
{"ilike", common.FilterOption{Column: "name", Operator: "ilike", Value: "%a%"}, `CAST("province_state"."name" AS TEXT) ILIKE ?`},
{"like", common.FilterOption{Column: "name", Operator: "like", Value: "%a%"}, `CAST("province_state"."name" AS TEXT) LIKE ?`},
{"in", common.FilterOption{Column: "name", Operator: "in", Value: []string{"a", "b"}}, `"province_state"."name" IN (?,?)`},
{"non-model column untouched", common.FilterOption{Column: "other", Operator: "eq", Value: 1}, `other = ?`},
{"already qualified untouched", common.FilterOption{Column: "rel.name", Operator: "eq", Value: 1}, `rel.name = ?`},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, _ := h.buildFilterConditionAlias(tt.filter, model, "province_state")
if got != tt.want {
t.Fatalf("got %q, want %q", got, tt.want)
}
})
}
}
func TestBuildFilterConditionAlias_NoAliasUnchanged(t *testing.T) {
h := &Handler{}
got, _ := h.buildFilterConditionAlias(common.FilterOption{Column: "name", Operator: "ilike", Value: "%a%"}, jsonColModel{}, "")
if got != "CAST(name AS TEXT) ILIKE ?" {
t.Fatalf("got %q", got)
}
}
func TestApplyFilter_QualifiesWithAlias(t *testing.T) {
h := &Handler{}
q := &jsonCapQuery{}
h.applyFilter(q, common.FilterOption{Column: "name", Operator: "ilike", Value: "%a%"}, jsonColModel{}, "province_state")
c := q.only(t)
if c.query != `CAST("province_state"."name" AS TEXT) ILIKE ?` {
t.Fatalf("query = %q", c.query)
}
}
func TestApplyFilters_OrGroupQualified(t *testing.T) {
h := &Handler{}
q := &jsonCapQuery{}
h.applyFilters(q, []common.FilterOption{
{Column: "name", Operator: "ilike", Value: "%a%"},
{Column: "id", Operator: "eq", Value: 1, LogicOperator: "OR"},
}, jsonColModel{}, "t")
c := q.only(t)
want := `(CAST("t"."name" AS TEXT) ILIKE ? OR "t"."id" = ?)`
if c.query != want {
t.Fatalf("query = %q, want %q", c.query, want)
}
}
+30 -15
View File
@@ -172,6 +172,9 @@ func (h *Handler) Handle(w common.ResponseWriter, r common.Request, params map[s
// Add request-scoped data to context
ctx = WithRequestData(ctx, schema, entity, tableName, model, modelPtr)
// Accept "<main table or alias>.<column>" for model columns before validation drops them
common.NormalizeMainTableFilters(model, tableName, &req.Options)
// Validate and filter columns in options (log warnings for invalid columns)
validator := common.NewColumnValidator(model)
req.Options = validator.FilterRequestOptions(req.Options)
@@ -411,7 +414,7 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
}
// Apply filters with proper grouping for OR logic
query = h.applyFilters(query, options.Filters, model)
query = h.applyFilters(query, options.Filters, model, common.MainTableAlias(model, tableName))
// Apply custom operators
for _, customOp := range options.CustomOperators {
@@ -558,7 +561,7 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
// Apply the same filters as the main query
for _, filter := range options.Filters {
rowNumQuery = h.applyFilter(rowNumQuery, filter, model)
rowNumQuery = h.applyFilter(rowNumQuery, filter, model, common.MainTableAlias(model, tableName))
}
// Apply custom operators
@@ -1932,7 +1935,7 @@ func (h *Handler) executeDelete(ctx context.Context, tx common.Database, hookCtx
// applyFilters applies all filters with proper grouping for OR logic
// Groups consecutive OR filters together to ensure proper query precedence
// Example: [A, B(OR), C(OR), D(AND)] => WHERE (A OR B OR C) AND D
func (h *Handler) applyFilters(query common.SelectQuery, filters []common.FilterOption, model interface{}) common.SelectQuery {
func (h *Handler) applyFilters(query common.SelectQuery, filters []common.FilterOption, model interface{}, alias string) common.SelectQuery {
if len(filters) == 0 {
return query
}
@@ -1952,11 +1955,11 @@ func (h *Handler) applyFilters(query common.SelectQuery, filters []common.Filter
}
// Apply the OR group as a single grouped WHERE clause
query = h.applyFilterGroup(query, orGroup, model)
query = h.applyFilterGroup(query, orGroup, model, alias)
i = j
} else {
// Single filter with AND logic (or first filter)
condition, args := h.buildFilterCondition(filters[i], model)
condition, args := h.buildFilterConditionAlias(filters[i], model, alias)
if condition != "" {
query = query.Where(condition, args...)
}
@@ -1969,7 +1972,7 @@ func (h *Handler) applyFilters(query common.SelectQuery, filters []common.Filter
// applyFilterGroup applies a group of filters that should be OR'd together
// Always wraps them in parentheses and applies as a single WHERE clause
func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.FilterOption, model interface{}) common.SelectQuery {
func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.FilterOption, model interface{}, alias string) common.SelectQuery {
if len(filters) == 0 {
return query
}
@@ -1979,7 +1982,7 @@ func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.Fi
var args []interface{}
for _, filter := range filters {
condition, filterArgs := h.buildFilterCondition(filter, model)
condition, filterArgs := h.buildFilterConditionAlias(filter, model, alias)
if condition != "" {
conditions = append(conditions, condition)
args = append(args, filterArgs...)
@@ -2005,6 +2008,12 @@ func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.Fi
// or the dotted data.x shorthand for a JSON column) resolve to a safe,
// parameterised expression before the ordinary operator handling below.
func (h *Handler) buildFilterCondition(filter common.FilterOption, model interface{}) (conditionString string, conditionArgs []interface{}) {
return h.buildFilterConditionAlias(filter, model, "")
}
// buildFilterConditionAlias is buildFilterCondition with plain model columns
// qualified by the main table alias, so joins from preloads can't make them ambiguous.
func (h *Handler) buildFilterConditionAlias(filter common.FilterOption, model interface{}, alias string) (conditionString string, conditionArgs []interface{}) {
var condition string
var args []interface{}
@@ -2012,6 +2021,9 @@ func (h *Handler) buildFilterCondition(filter common.FilterOption, model interfa
return cond, jargs
}
rawColumn := filter.Column
filter.Column = common.QualifyModelColumn(model, alias, filter.Column)
switch filter.Operator {
case "eq", "=":
condition = fmt.Sprintf("%s = ?", filter.Column)
@@ -2032,10 +2044,10 @@ func (h *Handler) buildFilterCondition(filter common.FilterOption, model interfa
condition = fmt.Sprintf("%s <= ?", filter.Column)
args = []interface{}{filter.Value}
case "like":
condition = fmt.Sprintf("%s LIKE ?", likeColumn(filter.Column, model))
condition = fmt.Sprintf("%s LIKE ?", likeColumn(filter.Column, rawColumn, model))
args = []interface{}{filter.Value}
case "ilike":
condition = fmt.Sprintf("%s ILIKE ?", likeColumn(filter.Column, model))
condition = fmt.Sprintf("%s ILIKE ?", likeColumn(filter.Column, rawColumn, model))
args = []interface{}{filter.Value}
case "in":
condition, args = common.BuildInCondition(filter.Column, filter.Value)
@@ -2073,14 +2085,14 @@ func (h *Handler) buildFilterCondition(filter common.FilterOption, model interfa
// CAST(... AS TEXT) would switch to case-sensitive matching and defeat a
// citext index. Every other column is cast to TEXT so LIKE/ILIKE also works
// against date/time/timestamp and numeric columns.
func likeColumn(column string, model interface{}) string {
if reflection.IsCitextColumn(model, column) {
func likeColumn(column, rawColumn string, model interface{}) string {
if reflection.IsCitextColumn(model, rawColumn) {
return column
}
return fmt.Sprintf("CAST(%s AS TEXT)", column)
}
func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOption, model interface{}) common.SelectQuery {
func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOption, model interface{}, alias string) common.SelectQuery {
// Determine which method to use based on LogicOperator
useOrLogic := strings.EqualFold(filter.LogicOperator, "OR")
@@ -2094,6 +2106,9 @@ func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOpti
return query.Where(cond, jargs...)
}
rawColumn := filter.Column
filter.Column = common.QualifyModelColumn(model, alias, filter.Column)
switch filter.Operator {
case "eq", "=":
condition = fmt.Sprintf("%s = ?", filter.Column)
@@ -2114,10 +2129,10 @@ func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOpti
condition = fmt.Sprintf("%s <= ?", filter.Column)
args = []interface{}{filter.Value}
case "like":
condition = fmt.Sprintf("%s LIKE ?", likeColumn(filter.Column, model))
condition = fmt.Sprintf("%s LIKE ?", likeColumn(filter.Column, rawColumn, model))
args = []interface{}{filter.Value}
case "ilike":
condition = fmt.Sprintf("%s ILIKE ?", likeColumn(filter.Column, model))
condition = fmt.Sprintf("%s ILIKE ?", likeColumn(filter.Column, rawColumn, model))
args = []interface{}{filter.Value}
case "in":
condition, args = common.BuildInCondition(filter.Column, filter.Value)
@@ -2523,7 +2538,7 @@ func (h *Handler) applyPreloads(model interface{}, query common.SelectQuery, pre
if len(preload.Filters) > 0 {
for _, filter := range preload.Filters {
sq = h.applyFilter(sq, filter, nil)
sq = h.applyFilter(sq, filter, nil, "")
}
}
if len(preload.Sort) > 0 {
+15 -14
View File
@@ -1,3 +1,4 @@
//go:build integration
// +build integration
package resolvespec
@@ -22,12 +23,12 @@ import (
// Test models
type TestUser struct {
ID uint `gorm:"primaryKey" json:"id"`
Name string `gorm:"not null" json:"name"`
Email string `gorm:"uniqueIndex;not null" json:"email"`
Age int `json:"age"`
Active bool `gorm:"default:true" json:"active"`
CreatedAt time.Time `json:"created_at"`
ID uint `gorm:"primaryKey" json:"id"`
Name string `gorm:"not null" json:"name"`
Email string `gorm:"uniqueIndex;not null" json:"email"`
Age int `json:"age"`
Active bool `gorm:"default:true" json:"active"`
CreatedAt time.Time `json:"created_at"`
Posts []TestPost `gorm:"foreignKey:UserID" json:"posts,omitempty"`
}
@@ -36,13 +37,13 @@ func (TestUser) TableName() string {
}
type TestPost struct {
ID uint `gorm:"primaryKey" json:"id"`
UserID uint `gorm:"not null" json:"user_id"`
Title string `gorm:"not null" json:"title"`
Content string `json:"content"`
Published bool `gorm:"default:false" json:"published"`
CreatedAt time.Time `json:"created_at"`
User *TestUser `gorm:"foreignKey:UserID" json:"user,omitempty"`
ID uint `gorm:"primaryKey" json:"id"`
UserID uint `gorm:"not null" json:"user_id"`
Title string `gorm:"not null" json:"title"`
Content string `json:"content"`
Published bool `gorm:"default:false" json:"published"`
CreatedAt time.Time `json:"created_at"`
User *TestUser `gorm:"foreignKey:UserID" json:"user,omitempty"`
Comments []TestComment `gorm:"foreignKey:PostID" json:"comments,omitempty"`
}
@@ -55,7 +56,7 @@ type TestComment struct {
PostID uint `gorm:"not null" json:"post_id"`
Content string `gorm:"not null" json:"content"`
CreatedAt time.Time `json:"created_at"`
Post *TestPost `gorm:"foreignKey:PostID" json:"post,omitempty"`
Post *TestPost `gorm:"foreignKey:PostID" json:"post,omitempty"`
}
func (TestComment) TableName() string {
+1 -1
View File
@@ -142,7 +142,7 @@ func TestApplyFilter_JSONColumn(t *testing.T) {
q := &jsonCapQuery{}
h.applyFilter(q, common.FilterOption{
Column: "data->>'tier'", Operator: "in", Value: []string{"a", "b"}, LogicOperator: "OR",
}, model)
}, model, "")
c := q.only(t)
if c.method != "WhereOr" || c.query != `("data" #>> ?::text[]) IN (?,?)` {
t.Fatalf("call = %+v", c)
+180
View File
@@ -0,0 +1,180 @@
package restheadspec
import (
"context"
"database/sql"
"testing"
"github.com/uptrace/bun"
"github.com/uptrace/bun/dialect/sqlitedialect"
"github.com/uptrace/bun/driver/sqliteshim"
"github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
)
type joinCountry struct {
bun.BaseModel `bun:"table:country,alias:country"`
ID int64 `bun:"id,pk"`
Name string `bun:"name"`
}
type joinProvince struct {
bun.BaseModel `bun:"table:province_state,alias:province_state"`
ID int64 `bun:"id,pk"`
Name string `bun:"name"`
Abbreviation string `bun:"abbreviation"`
RidCountry int64 `bun:"rid_country"`
}
func setupJoinDB(t *testing.T) *bun.DB {
t.Helper()
sqldb, err := sql.Open(sqliteshim.ShimName, "file:rhsfilterjoin?mode=memory&cache=private")
if err != nil {
t.Fatal(err)
}
db := bun.NewDB(sqldb, sqlitedialect.New())
t.Cleanup(func() { _ = db.Close() })
ctx := context.Background()
for _, m := range []interface{}{(*joinCountry)(nil), (*joinProvince)(nil)} {
if _, err := db.NewCreateTable().Model(m).IfNotExists().Exec(ctx); err != nil {
t.Fatal(err)
}
}
if _, err := db.NewInsert().Model(&[]joinCountry{{ID: 1, Name: "Abcland"}, {ID: 2, Name: "Other"}}).Exec(ctx); err != nil {
t.Fatal(err)
}
if _, err := db.NewInsert().Model(&[]joinProvince{
{ID: 1, Name: "abc one", Abbreviation: "A1", RidCountry: 1},
{ID: 2, Name: "xyz two", Abbreviation: "ABC", RidCountry: 1},
{ID: 3, Name: "nope", Abbreviation: "N3", RidCountry: 2},
}).Exec(ctx); err != nil {
t.Fatal(err)
}
return db
}
const joinSQL = "LEFT JOIN country AS rel_rid_country ON rel_rid_country.id = province_state.rid_country"
// countFilters applies filters the way handleRead does (single AND filters via
// applyFilter, consecutive OR filters via applyOrFilterGroup) and counts.
func countFilters(t *testing.T, db *bun.DB, withJoin bool, filters []common.FilterOption) (int, error) {
t.Helper()
h := &Handler{}
model := &joinProvince{}
var q common.SelectQuery = database.NewBunAdapter(db).NewSelect().Model(&[]*joinProvince{})
if withJoin {
q = q.Join(joinSQL)
}
for i := 0; i < len(filters); {
f := filters[i]
castInfo := h.ValidateAndAdjustFilterForColumnType(&f, model)
if f.LogicOperator == "OR" {
group := []*common.FilterOption{&f}
info := []ColumnCastInfo{castInfo}
j := i + 1
for j < len(filters) && filters[j].LogicOperator == "OR" {
g := filters[j]
info = append(info, h.ValidateAndAdjustFilterForColumnType(&g, model))
group = append(group, &g)
j++
}
q = h.applyOrFilterGroup(q, group, info, "public.province_state", model)
i = j
continue
}
q = h.applyFilter(q, f, "public.province_state", castInfo.NeedsCast, "AND", model)
i++
}
return q.Count(context.Background())
}
func TestRHSFilters_WithAndWithoutJoin_RealQuery(t *testing.T) {
db := setupJoinDB(t)
cases := []struct {
name string
filters []common.FilterOption
want int
}{
{"eq", []common.FilterOption{{Column: "name", Operator: "eq", Value: "nope"}}, 1},
{"neq", []common.FilterOption{{Column: "name", Operator: "neq", Value: "nope"}}, 2},
{"like", []common.FilterOption{{Column: "name", Operator: "like", Value: "%abc%"}}, 1},
{"in", []common.FilterOption{{Column: "name", Operator: "in", Value: []string{"nope", "xyz two"}}}, 2},
{"gt", []common.FilterOption{{Column: "id", Operator: "gt", Value: 1}}, 2},
{"between", []common.FilterOption{{Column: "id", Operator: "between", Value: []interface{}{0, 3}}}, 2},
{"between_inclusive", []common.FilterOption{{Column: "id", Operator: "between_inclusive", Value: []interface{}{1, 3}}}, 3},
{"is_not_null", []common.FilterOption{{Column: "name", Operator: "is_not_null"}}, 3},
{"or group", []common.FilterOption{
{Column: "name", Operator: "like", Value: "%abc%", LogicOperator: "OR"},
{Column: "abbreviation", Operator: "like", Value: "%abc%", LogicOperator: "OR"},
}, 2},
{"or group then and", []common.FilterOption{
{Column: "name", Operator: "like", Value: "%abc%", LogicOperator: "OR"},
{Column: "abbreviation", Operator: "like", Value: "%abc%", LogicOperator: "OR"},
{Column: "id", Operator: "eq", Value: 2},
}, 1},
}
for _, c := range cases {
for _, withJoin := range []bool{false, true} {
name := c.name + "/no_join"
if withJoin {
name = c.name + "/join"
}
t.Run(name, func(t *testing.T) {
n, err := countFilters(t, db, withJoin, c.filters)
if err != nil {
t.Fatalf("query failed (ambiguous column?): %v", err)
}
if n != c.want {
t.Fatalf("count = %d, want %d", n, c.want)
}
})
}
}
}
func TestRHSFilters_JoinedTableColumnPassesThrough(t *testing.T) {
db := setupJoinDB(t)
n, err := countFilters(t, db, true, []common.FilterOption{
{Column: "rel_rid_country.name", Operator: "eq", Value: "Abcland"},
})
if err != nil {
t.Fatal(err)
}
if n != 2 {
t.Fatalf("count = %d, want 2", n)
}
}
// Client sends "province_state.name": it must survive validation and be applied.
func TestRHSFilters_ClientQualifiedColumn_NotDropped(t *testing.T) {
db := setupJoinDB(t)
h := &Handler{}
model := &joinProvince{}
opts := ExtendedRequestOptions{}
opts.Filters = []common.FilterOption{{Column: "province_state.name", Operator: "like", Value: "%abc%"}}
common.NormalizeMainTableFilters(model, "public.province_state", &opts.RequestOptions)
opts = h.filterExtendedOptions(common.NewColumnValidator(model), opts, model)
if len(opts.Filters) != 1 || opts.Filters[0].Column != "name" {
t.Fatalf("filter was dropped or not normalised: %+v", opts.Filters)
}
n, err := countFilters(t, db, true, opts.Filters)
if err != nil {
t.Fatal(err)
}
if n != 1 {
t.Fatalf("count = %d, want 1", n)
}
}
func TestRHSFilters_ILikeSQLQualified(t *testing.T) {
h := &Handler{}
q := &jsonCapQuery{}
f := common.FilterOption{Column: "name", Operator: "ilike", Value: "%abc%"}
h.applyFilter(q, f, "info.province_state", false, "AND", jsonColModel{})
if got := q.only(t).query; got != "CAST(province_state.name AS TEXT) ILIKE ?" {
t.Fatalf("query = %q", got)
}
}
+3
View File
@@ -164,6 +164,9 @@ func (h *Handler) Handle(w common.ResponseWriter, r common.Request, params map[s
// Parse options from headers - this now includes relation name resolution
options := h.parseOptionsFromHeaders(r, model)
// Accept "<main table or alias>.<column>" for model columns before validation drops them
common.NormalizeMainTableFilters(model, tableName, &options.RequestOptions)
// Validate and filter columns in options (log warnings for invalid columns)
validator := common.NewColumnValidator(model)
options = h.filterExtendedOptions(validator, options, model)