Compare commits

...
19 Commits
Author SHA1 Message Date
Hein a74eebc7f3 fix(json-columns): skip JSON select columns with no model scan target
Tests / Unit Tests (push) Failing after 28s
Tests / Integration Tests (push) Failing after 29s
Build , Vet Test, and Lint / Build (push) Successful in 1m12s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 1m43s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 1m46s
Build , Vet Test, and Lint / Lint Code (push) Successful in 1m48s
Requesting a JSON sub-field (e.g. jsonvalue->'product'->>'cost') that has
no matching bun scanonly field on the model made the whole read fail with
"bun: ModelX does not have column Y", since bun scans SELECT results
straight into the typed model struct.

Add reflection.HasColumn to check whether the model can actually receive
a given column (including scanonly fields, walking embedded structs), and
gate the JSON select-column expression on it in ApplySelectColumns
(shared by websocketspec/mqttspec) and the resolvespec/restheadspec
handlers. When there's no scan target, drop just that column with a
warning instead of erroring the whole request.
2026-09-28 12:30:56 +02:00
Hein 6687a7a5cd fix(bun): spread variadic args correctly in ColumnExpr 2026-09-28 12:01:38 +02:00
warkanum e8fbbede7e chore: update version to 1.0.2 and changelog
Tests / Unit Tests (push) Failing after 25s
Tests / Integration Tests (push) Failing after 26s
Build , Vet Test, and Lint / Build (push) Successful in 1m9s
Build , Vet Test, and Lint / Lint Code (push) Successful in 1m29s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 1m32s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 1m33s
2026-09-23 20:58:18 +02:00
warkanum 7f8982fa35 feat(headerspec): add UTF-8 encoding/decoding support
* Implement round-trip encoding/decoding for UTF-8 values in header functions.
* Update encodeHeaderValue and decodeHeaderValue to use base64 utility functions.
* Add tests for UTF-8 header value handling.
* Update Vite config to externalize base64 utility.
2026-09-23 20:57:50 +02:00
warkanum b587cbd3c4 feat(headers): add support for custom HTTP headers in clients
* Introduced `headers` property in `ClientConfig` interface.
* Updated `HeaderSpecClient` and `ResolveSpecClient` to utilize custom headers.
* Implemented `mergeHeaders` function to handle case-insensitive header merging.
* Added tests for custom header functionality in clients.
2026-09-23 19:47:05 +02:00
Hein a220338eea feat(spectypes): add CIString, LCString, and UCString types with tests
Tests / Unit Tests (push) Failing after 28s
Tests / Integration Tests (push) Failing after 30s
Build , Vet Test, and Lint / Build (push) Successful in 1m8s
Build , Vet Test, and Lint / Lint Code (push) Successful in 1m18s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 1m35s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 1m36s
2026-09-21 14:06:33 +02:00
Hein 20c67166d0 fix(handler): support implicit updates from request body 2026-09-21 11:18:10 +02:00
Hein 749dad4ed1 fix(quickproxy): ensure request body is preserved on fallback
Tests / Unit Tests (push) Failing after 26s
Tests / Integration Tests (push) Failing after 41s
Build , Vet Test, and Lint / Build (push) Successful in 4m26s
Build , Vet Test, and Lint / Lint Code (push) Successful in 4m58s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 5m1s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 5m3s
2026-09-21 09:21:58 +02:00
warkanum d6c5740f9c fix(handler): add operation type to hook context
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Failing after 1s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Failing after 1s
Build , Vet Test, and Lint / Lint Code (push) Failing after 1s
Build , Vet Test, and Lint / Build (push) Failing after 1s
Tests / Unit Tests (push) Failing after 0s
Tests / Integration Tests (push) Failing after 10s
2026-09-20 16:12:44 +02:00
warkanum 817b781c88 fix(security): skip loading security rules if disabled 2026-09-20 15:52:25 +02:00
warkanum 87eaa9e18c fix(security): skip row security enforcement for specific operations
* Add ShouldSkipRowSecurity function to determine when to bypass row security
* Update ApplyRowSecurity to utilize operation context for enforcement
2026-09-20 15:51:02 +02:00
Hein 4f6878099b fix(quickproxy): reject Exclude entries outside their rule's URLPrefix
Tests / Unit Tests (push) Failing after 5s
Tests / Integration Tests (push) Failing after 23s
Build , Vet Test, and Lint / Build (push) Successful in 52s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 55s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 56s
Build , Vet Test, and Lint / Lint Code (push) Successful in 1m4s
An Exclude entry only ever matches requests that already fall under
its rule's URLPrefix, so one written without that prefix (e.g.
"/health" on a rule for "/api") silently never triggered. Validate
that each Exclude entry itself starts with the rule's URLPrefix,
failing NewService instead of accepting a no-op config.
2026-09-17 11:40:49 +02:00
Hein 0d8b136b91 feat(quickproxy): support per-rule Exclude path prefixes
Rule.Exclude lists path prefixes that should never be proxied by that
rule, even though they fall under its URLPrefix. A request matching an
Exclude prefix is treated as a non-match for that rule: matching
continues against other configured rules, falling back to the
caller-supplied handler if none apply. Lets a catch-all "/" rule proxy
everything except carved-out paths like "/health".
2026-09-17 11:38:02 +02:00
Hein 6de9be0ae7 chore(proxy): remove outdated proxy documentation 2026-09-17 11:09:14 +02:00
Hein 82e923b16e fix(quickproxy): use http.NotFoundHandler per golangci-lint gocritic 2026-09-17 11:08:12 +02:00
Hein 9a664593f0 feat(quickproxy): add reverse-proxy-with-static-fallback package
Adds pkg/server/quickproxy: longest-prefix rule matching over
net/http/httputil.ReverseProxy, falling back to a caller-supplied
handler when the upstream is unreachable or returns 404. Any other
upstream response streams through unchanged. All HTTP methods are
proxied, with a configurable global dial/response-header timeout
(quickproxy.WithTimeout, default 10s).

GoCore-side wiring (config field, webserver2/proxy.go, server.go
route ordering) is tracked separately in that repo.
2026-09-17 11:07:53 +02:00
Hein 6e3124e4e0 fix(restheadspec): prevent casting numeric values in ILIKE filters
Tests / Unit Tests (push) Failing after 5s
Tests / Integration Tests (push) Failing after 24s
Build , Vet Test, and Lint / Build (push) Successful in 54s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 1m6s
Build , Vet Test, and Lint / Lint Code (push) Successful in 2m38s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 2m47s
2026-09-15 13:55:11 +02:00
Hein d5de48011b fix(websocketspec,mqttspec,resolvemcp): stop casting citext columns to TEXT for LIKE/ILIKE
Same class of bug as the restheadspec/resolvespec fix: these handlers
unconditionally rendered CAST(col AS TEXT) LIKE/ILIKE for every column,
which flips a citext column to case-sensitive matching and defeats a
citext index. Thread the model through to buildFilterCondition/applyFilters
so reflection.IsCitextColumn can skip the cast for citext columns.

resolvemcp's eq/neq/gt/lt paths never cast (they never had the
restheadspec-style reflect.Kind cast heuristic), so this only touches
LIKE/ILIKE. funcspec is unaffected: it has no Go struct model to check
against (colname/value come straight from SQL function parameters).
2026-09-15 11:26:28 +02:00
Hein 6bd6a6f164 fix(restheadspec,resolvespec): stop casting SqlNull-wrapped and citext columns to TEXT in filters
reflect.Type.Kind() on spectypes.SqlNull[T] wrappers (SqlInt16/32/64,
SqlFloat64, SqlBool, SqlString, and embedders like SqlTimeStamp) always
reports reflect.Struct, never the wrapped T. ValidateAndAdjustFilterForColumnType
treated those as "complex" columns and forced CAST(col AS TEXT) on eq/gt/lt
filters, e.g. CAST(atdetail.rid_parent AS TEXT) = '90446096', which can't
use the index on rid_parent.

Add spectypes.UnwrapKind to see through SqlNull wrappers to the underlying
Kind, and use it in GetColumnTypeFromModel so numeric/string SqlNull columns
are recognized correctly and compared natively.

Also stop unconditionally casting to TEXT for LIKE/ILIKE and add
reflection.IsCitextColumn: citext columns are already case-insensitive, so
casting them to TEXT flips to case-sensitive matching and defeats a citext
index.
2026-09-15 11:18:13 +02:00
36 changed files with 3834 additions and 2323 deletions
+1 -1
View File
@@ -339,7 +339,7 @@ func (b *BunSelectQuery) Column(columns ...string) common.SelectQuery {
func (b *BunSelectQuery) ColumnExpr(query string, args ...interface{}) common.SelectQuery {
if len(args) > 0 {
b.query = b.query.ColumnExpr(query, args)
b.query = b.query.ColumnExpr(query, args...)
} else {
b.query = b.query.ColumnExpr(query)
}
@@ -0,0 +1,31 @@
package database
import (
"strings"
"testing"
"github.com/stretchr/testify/require"
)
// TestBunSelectQuery_ColumnExpr_SpreadsArgs is a regression test for a bug
// where ColumnExpr passed its variadic args slice as a single argument
// (b.query.ColumnExpr(query, args) instead of args...), causing bun to
// serialize the arg slice itself (e.g. producing `'["{product,cost}"]'`
// instead of `'{product,cost}'` for a JSON path parameter).
func TestBunSelectQuery_ColumnExpr_SpreadsArgs(t *testing.T) {
db := setupBunTestDB(t)
defer db.Close()
adapter := NewBunAdapter(db)
sq := adapter.NewSelect().
Table("test_inserts").
ColumnExpr("(jsonvalue #>> ?::text[]) AS jsonvalue_product_cost", "{product,cost}")
bsq, ok := sq.(*BunSelectQuery)
require.True(t, ok, "expected *BunSelectQuery")
sqlStr := bsq.query.String()
require.NotContains(t, sqlStr, `["{product,cost}"]`, "arg slice must not be serialized as a JSON array: %s", sqlStr)
require.True(t, strings.Contains(sqlStr, `'{product,cost}'`), "expected the bound text[] literal in SQL: %s", sqlStr)
}
+10
View File
@@ -4,6 +4,7 @@ import (
"fmt"
"strings"
"github.com/bitechdev/ResolveSpec/pkg/logger"
"github.com/bitechdev/ResolveSpec/pkg/reflection"
)
@@ -72,6 +73,15 @@ func ResolveJSONColumnExpr(model interface{}, tableAlias, token string) (expr st
func ApplySelectColumns(query SelectQuery, model interface{}, tableAlias string, columns []string) SelectQuery {
for _, col := range columns {
if expr, args, alias, ok := ResolveJSONColumnExpr(model, tableAlias, col); ok {
if !reflection.HasColumn(model, alias) {
// No matching scan target on the model (e.g. no
// `bun:"<alias>,scanonly"` field declared for this JSON
// path) - bun would fail to scan the row with "does not
// have column X". Drop the expression rather than erroring;
// the rest of the requested columns still get selected.
logger.Warn("Skipping JSON select column %q: model has no scan target for alias %q", col, alias)
continue
}
query = query.ColumnExpr(expr+" AS "+QuoteIdent(alias), args...)
continue
}
+41
View File
@@ -2,6 +2,7 @@ package common
import (
"reflect"
"strings"
"testing"
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
@@ -162,3 +163,43 @@ func TestBuildJSONFilterCondition_QualifiedAndInjectionSafe(t *testing.T) {
t.Errorf("args = %#v", args)
}
}
// selectCapQuery is a minimal SelectQuery that records Column/ColumnExpr calls
// so ApplySelectColumns' behaviour can be asserted without a real DB.
type selectCapQuery struct {
SelectQuery
columns []string
columnExprs []string
}
func (m *selectCapQuery) Column(cols ...string) SelectQuery {
m.columns = append(m.columns, cols...)
return m
}
func (m *selectCapQuery) ColumnExpr(q string, args ...interface{}) SelectQuery {
m.columnExprs = append(m.columnExprs, q)
return m
}
// jsonSelectModel has a real JSON column (Data) but only ONE pre-declared
// scanonly field for a computed JSON path ("data_city"); "data_age" has no
// matching scan target.
type jsonSelectModel struct {
ID int64 `json:"id" bun:"id,pk"`
Data spectypes.SqlJSONB `json:"data" bun:"data"`
DataCity string `json:"-" bun:"data_city,scanonly"`
}
func TestApplySelectColumns_SkipsJSONColumnWithoutScanTarget(t *testing.T) {
m := jsonSelectModel{}
q := &selectCapQuery{}
ApplySelectColumns(q, m, "", []string{"id", "data.city", "data.age"})
if !reflect.DeepEqual(q.columns, []string{"id"}) {
t.Errorf("columns = %#v, want [id]", q.columns)
}
if len(q.columnExprs) != 1 || !strings.Contains(q.columnExprs[0], `AS "data_city"`) {
t.Errorf("columnExprs = %#v, want exactly one expr aliased data_city", q.columnExprs)
}
}
+14 -2
View File
@@ -720,7 +720,13 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata
}
op := strings.ToLower(filter.Operator)
if op == "like" || op == "ilike" {
query = query.Where(fmt.Sprintf("CAST(%s AS TEXT) %s ?", filter.Column, h.getOperatorSQL(filter.Operator)), filter.Value)
// citext columns are already case-insensitive; casting to TEXT would
// switch to case-sensitive matching and defeat a citext index.
if reflection.IsCitextColumn(hookCtx.Model, filter.Column) {
query = query.Where(fmt.Sprintf("%s %s ?", filter.Column, h.getOperatorSQL(filter.Operator)), filter.Value)
} else {
query = query.Where(fmt.Sprintf("CAST(%s AS TEXT) %s ?", filter.Column, h.getOperatorSQL(filter.Operator)), filter.Value)
}
} else {
query = query.Where(fmt.Sprintf("%s %s ?", filter.Column, h.getOperatorSQL(filter.Operator)), filter.Value)
}
@@ -786,7 +792,13 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata
}
op := strings.ToLower(filter.Operator)
if op == "like" || op == "ilike" {
countQuery = countQuery.Where(fmt.Sprintf("CAST(%s AS TEXT) %s ?", filter.Column, h.getOperatorSQL(filter.Operator)), filter.Value)
// citext columns are already case-insensitive; casting to TEXT would
// switch to case-sensitive matching and defeat a citext index.
if reflection.IsCitextColumn(hookCtx.Model, filter.Column) {
countQuery = countQuery.Where(fmt.Sprintf("%s %s ?", filter.Column, h.getOperatorSQL(filter.Operator)), filter.Value)
} else {
countQuery = countQuery.Where(fmt.Sprintf("CAST(%s AS TEXT) %s ?", filter.Column, h.getOperatorSQL(filter.Operator)), filter.Value)
}
} else {
countQuery = countQuery.Where(fmt.Sprintf("%s %s ?", filter.Column, h.getOperatorSQL(filter.Operator)), filter.Value)
}
+61 -3
View File
@@ -9,6 +9,7 @@ import (
"time"
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
)
type PrimaryKeyNameProvider interface {
@@ -438,6 +439,63 @@ func GetSQLModelColumns(model any) []string {
return columns
}
// HasColumn reports whether the model has a struct field that bun/gorm would
// scan a column named columnName into. Unlike GetSQLModelColumns, this
// includes scanonly fields (e.g. a `bun:"jsonvalue_product_cost,scanonly"`
// field added specifically to receive a computed/JSON-path SELECT expression)
// since those are legitimate scan targets even though they are not writable.
// Matching is case-insensitive against the resolved bun/gorm/json column name
// and against the bare Go field name.
func HasColumn(model any, columnName string) bool {
if columnName == "" {
return false
}
modelType := reflect.TypeOf(model)
for modelType != nil && (modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array) {
modelType = modelType.Elem()
}
if modelType == nil || modelType.Kind() != reflect.Struct {
return false
}
return hasColumnInType(modelType, columnName)
}
func hasColumnInType(typ reflect.Type, columnName string) bool {
for i := 0; i < typ.NumField(); i++ {
field := typ.Field(i)
if !field.IsExported() {
continue
}
bunTag := field.Tag.Get("bun")
gormTag := field.Tag.Get("gorm")
if field.Anonymous {
fieldType := field.Type
if fieldType.Kind() == reflect.Pointer {
fieldType = fieldType.Elem()
}
if fieldType.Kind() == reflect.Struct {
if hasColumnInType(fieldType, columnName) {
return true
}
continue
}
}
if bunTag == "-" || gormTag == "-" {
continue
}
if strings.EqualFold(getColumnNameFromField(field), columnName) || strings.EqualFold(field.Name, columnName) {
return true
}
}
return false
}
// collectSQLColumnsFromType recursively collects SQL column names from a struct type
// scanOnlyEmbedded indicates if we're inside a scan-only embedded struct
func collectSQLColumnsFromType(typ reflect.Type, columns *[]string, scanOnlyEmbedded bool) {
@@ -728,19 +786,19 @@ func GetColumnTypeFromModel(model interface{}, colName string) reflect.Kind {
// Parse JSON tag (format: "name,omitempty")
parts := strings.Split(jsonTag, ",")
if parts[0] == sourceColName {
return field.Type.Kind()
return spectypes.UnwrapKind(field.Type)
}
}
// Check field name (case-insensitive)
if strings.EqualFold(field.Name, sourceColName) {
return field.Type.Kind()
return spectypes.UnwrapKind(field.Type)
}
// Check snake_case conversion
snakeCaseName := ToSnakeCase(field.Name)
if snakeCaseName == sourceColName {
return field.Type.Kind()
return spectypes.UnwrapKind(field.Type)
}
}
+35
View File
@@ -3,6 +3,8 @@ package reflection
import (
"reflect"
"testing"
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
)
// Test models for GORM
@@ -409,6 +411,23 @@ func TestGetModelColumnsWithEmbedded(t *testing.T) {
}
}
func TestHasColumn(t *testing.T) {
m := ModelWithEmbedded{}
for _, col := range []string{"name", "description", "rid_base", "created_at", "cql1", "cql2"} {
if !HasColumn(m, col) {
t.Errorf("HasColumn(%q) = false, want true", col)
}
}
if HasColumn(m, "nonexistent_column") {
t.Error("HasColumn(nonexistent_column) = true, want false")
}
if HasColumn(m, "") {
t.Error("HasColumn(\"\") = true, want false")
}
}
func TestIsColumnWritableWithEmbedded(t *testing.T) {
tests := []struct {
name string
@@ -1047,6 +1066,22 @@ func TestGetColumnTypeFromModel(t *testing.T) {
}
}
// SqlNull-wrapped columns (e.g. nullable bigint foreign keys) must report the
// wrapped value's Kind, not reflect.Struct, so numeric eq/gt/lt filters don't
// get an unnecessary CAST(... AS TEXT) that defeats the column's index.
type SqlNullFKModel struct {
RidParent spectypes.SqlInt64 `bun:"rid_parent" json:"rid_parent"`
}
func TestGetColumnTypeFromModel_SqlNullWrapper(t *testing.T) {
model := SqlNullFKModel{RidParent: spectypes.NewSqlInt64(90446096)}
result := GetColumnTypeFromModel(model, "rid_parent")
if result != reflect.Int64 {
t.Errorf("GetColumnTypeFromModel(rid_parent) = %v, want %v (SqlInt64 must unwrap to its numeric Kind)", result, reflect.Int64)
}
}
// ============= Tests for relation functions =============
// Models for relation testing
+26 -8
View File
@@ -145,8 +145,16 @@ func IsJSONColumn(model interface{}, colName string) bool {
// tagDeclaresJSON reports whether an ORM struct tag declares a json/jsonb column
// type, e.g. `bun:"meta,type:jsonb"` or `gorm:"column:meta;type:json"`.
func tagDeclaresJSON(tag string) bool {
return columnTypeTagValue(tag) == "json" || strings.HasPrefix(columnTypeTagValue(tag), "json(") ||
columnTypeTagValue(tag) == "jsonb" || strings.HasPrefix(columnTypeTagValue(tag), "jsonb(")
}
// columnTypeTagValue extracts the lower-cased value of a `type:` entry from a
// bun or gorm struct tag, e.g. `bun:"name,type:citext"` -> "citext". Returns ""
// if the tag carries no `type:` entry.
func columnTypeTagValue(tag string) string {
if tag == "" {
return false
return ""
}
for _, part := range strings.FieldsFunc(tag, func(r rune) bool {
return r == ',' || r == ';' || r == ' '
@@ -155,12 +163,22 @@ func tagDeclaresJSON(tag string) bool {
if !found {
continue
}
value = strings.ToLower(strings.TrimSpace(value))
// Match "json" and "jsonb", including parametrised forms just in case.
if value == "json" || value == "jsonb" ||
strings.HasPrefix(value, "json(") || strings.HasPrefix(value, "jsonb(") {
return true
}
return strings.ToLower(strings.TrimSpace(value))
}
return false
return ""
}
// IsCitextColumn reports whether colName carries an explicit `type:citext`
// bun/gorm tag. citext columns must never be CAST(... AS TEXT) for comparisons:
// that swaps in case-sensitive text semantics and defeats any citext index.
func IsCitextColumn(model interface{}, colName string) bool {
f, ok := getColumnStructField(model, colName)
if !ok {
return false
}
tagVal := columnTypeTagValue(f.Tag.Get("bun"))
if tagVal == "" {
tagVal = columnTypeTagValue(f.Tag.Get("gorm"))
}
return tagVal == "citext"
}
+19 -10
View File
@@ -269,7 +269,7 @@ func (h *Handler) executeRead(ctx context.Context, schema, entity, id string, op
}
// Filters
query = h.applyFilters(query, options.Filters)
query = h.applyFilters(query, options.Filters, model)
// Custom operators
for _, customOp := range options.CustomOperators {
@@ -751,8 +751,10 @@ func (h *Handler) executeDelete(ctx context.Context, schema, entity, id string)
return recordToDelete, nil
}
// applyFilters applies all filters with OR grouping logic.
func (h *Handler) applyFilters(query common.SelectQuery, filters []common.FilterOption) common.SelectQuery {
// applyFilters applies all filters with OR grouping logic. model, when
// non-nil, lets citext columns be recognised so LIKE/ILIKE compares them
// natively instead of casting to TEXT (which would defeat a citext index).
func (h *Handler) applyFilters(query common.SelectQuery, filters []common.FilterOption, model interface{}) common.SelectQuery {
if len(filters) == 0 {
return query
}
@@ -768,10 +770,10 @@ func (h *Handler) applyFilters(query common.SelectQuery, filters []common.Filter
orGroup = append(orGroup, filters[j])
j++
}
query = h.applyFilterGroup(query, orGroup)
query = h.applyFilterGroup(query, orGroup, model)
i = j
} else {
condition, args := h.buildFilterCondition(filters[i])
condition, args := h.buildFilterCondition(filters[i], model)
if condition != "" {
query = query.Where(condition, args...)
}
@@ -782,12 +784,12 @@ func (h *Handler) applyFilters(query common.SelectQuery, filters []common.Filter
return query
}
func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.FilterOption) common.SelectQuery {
func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.FilterOption, model interface{}) common.SelectQuery {
var conditions []string
var args []interface{}
for _, filter := range filters {
condition, filterArgs := h.buildFilterCondition(filter)
condition, filterArgs := h.buildFilterCondition(filter, model)
if condition != "" {
conditions = append(conditions, condition)
args = append(args, filterArgs...)
@@ -803,7 +805,14 @@ func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.Fi
return query.Where("("+strings.Join(conditions, " OR ")+")", args...)
}
func (h *Handler) buildFilterCondition(filter common.FilterOption) (condition string, args []interface{}) {
func (h *Handler) buildFilterCondition(filter common.FilterOption, model interface{}) (condition string, args []interface{}) {
// citext columns are already case-insensitive; casting to TEXT would
// switch to case-sensitive matching and defeat a citext index.
likeColumn := filter.Column
if !reflection.IsCitextColumn(model, filter.Column) {
likeColumn = fmt.Sprintf("CAST(%s AS TEXT)", filter.Column)
}
switch filter.Operator {
case "eq", "=":
return fmt.Sprintf("%s = ?", filter.Column), []interface{}{filter.Value}
@@ -818,9 +827,9 @@ func (h *Handler) buildFilterCondition(filter common.FilterOption) (condition st
case "lte", "<=":
return fmt.Sprintf("%s <= ?", filter.Column), []interface{}{filter.Value}
case "like":
return fmt.Sprintf("CAST(%s AS TEXT) LIKE ?", filter.Column), []interface{}{filter.Value}
return fmt.Sprintf("%s LIKE ?", likeColumn), []interface{}{filter.Value}
case "ilike":
return fmt.Sprintf("CAST(%s AS TEXT) ILIKE ?", filter.Column), []interface{}{filter.Value}
return fmt.Sprintf("%s ILIKE ?", likeColumn), []interface{}{filter.Value}
case "in":
condition, args := common.BuildInCondition(filter.Column, filter.Value)
return condition, args
+133 -106
View File
@@ -306,15 +306,16 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
txErr := h.db.RunInTransaction(ctx, func(tx common.Database) error {
hookCtx := &HookContext{
Context: ctx,
Handler: h,
Schema: schema,
Entity: entity,
Model: model,
Options: options,
ID: id,
Writer: w,
Tx: tx,
Context: ctx,
Handler: h,
Schema: schema,
Entity: entity,
Model: model,
Operation: "read",
Options: options,
ID: id,
Writer: w,
Tx: tx,
}
if err := h.hooks.ExecuteBeforeOp(BeforeRead, hookCtx); err != nil {
statusCode, errCode, errMsg = http.StatusInternalServerError, "hook_error", "BeforeRead hook failed"
@@ -348,6 +349,10 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
logger.Debug("Selecting columns: %v", options.Columns)
for _, col := range options.Columns {
if expr, jargs, alias, ok := common.ResolveJSONColumnExpr(model, "", col); ok {
if !reflection.HasColumn(model, alias) {
logger.Warn("Skipping JSON select column %q: model has no scan target for alias %q", col, alias)
continue
}
query = query.ColumnExpr(expr+" AS "+common.QuoteIdent(alias), jargs...)
continue
}
@@ -722,15 +727,16 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
var nestedResult *common.ProcessResult
err := h.db.RunInTransaction(ctx, func(tx common.Database) error {
hookCtx := &HookContext{
Context: ctx,
Handler: h,
Schema: schema,
Entity: entity,
Model: model,
Options: options,
Data: v,
Writer: w,
Tx: tx,
Context: ctx,
Handler: h,
Schema: schema,
Entity: entity,
Model: model,
Operation: "create",
Options: options,
Data: v,
Writer: w,
Tx: tx,
}
if err := h.hooks.ExecuteBeforeOp(BeforeCreate, hookCtx); err != nil {
return fmt.Errorf("BeforeCreate hook failed: %w", err)
@@ -769,15 +775,16 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
var responseData interface{} = v
err := h.db.RunInTransaction(ctx, func(tx common.Database) error {
hookCtx := &HookContext{
Context: ctx,
Handler: h,
Schema: schema,
Entity: entity,
Model: model,
Options: options,
Data: v,
Writer: w,
Tx: tx,
Context: ctx,
Handler: h,
Schema: schema,
Entity: entity,
Model: model,
Operation: "create",
Options: options,
Data: v,
Writer: w,
Tx: tx,
}
if err := h.hooks.ExecuteBeforeOp(BeforeCreate, hookCtx); err != nil {
return fmt.Errorf("BeforeCreate hook failed: %w", err)
@@ -851,15 +858,16 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
for _, item := range v {
hookCtx := &HookContext{
Context: ctx,
Handler: h,
Schema: schema,
Entity: entity,
Model: model,
Options: options,
Data: item,
Writer: w,
Tx: tx,
Context: ctx,
Handler: h,
Schema: schema,
Entity: entity,
Model: model,
Operation: "create",
Options: options,
Data: item,
Writer: w,
Tx: tx,
}
if err := h.hooks.ExecuteBeforeOp(BeforeCreate, hookCtx); err != nil {
return fmt.Errorf("BeforeCreate hook failed: %w", err)
@@ -898,15 +906,16 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
err := h.db.RunInTransaction(ctx, func(tx common.Database) error {
for _, item := range v {
hookCtx := &HookContext{
Context: ctx,
Handler: h,
Schema: schema,
Entity: entity,
Model: model,
Options: options,
Data: item,
Writer: w,
Tx: tx,
Context: ctx,
Handler: h,
Schema: schema,
Entity: entity,
Model: model,
Operation: "create",
Options: options,
Data: item,
Writer: w,
Tx: tx,
}
if err := h.hooks.ExecuteBeforeOp(BeforeCreate, hookCtx); err != nil {
return fmt.Errorf("BeforeCreate hook failed: %w", err)
@@ -982,15 +991,16 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
for _, item := range v {
if itemMap, ok := item.(map[string]interface{}); ok {
hookCtx := &HookContext{
Context: ctx,
Handler: h,
Schema: schema,
Entity: entity,
Model: model,
Options: options,
Data: itemMap,
Writer: w,
Tx: tx,
Context: ctx,
Handler: h,
Schema: schema,
Entity: entity,
Model: model,
Operation: "create",
Options: options,
Data: itemMap,
Writer: w,
Tx: tx,
}
if err := h.hooks.ExecuteBeforeOp(BeforeCreate, hookCtx); err != nil {
return fmt.Errorf("BeforeCreate hook failed: %w", err)
@@ -1035,15 +1045,16 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
}
hookCtx := &HookContext{
Context: ctx,
Handler: h,
Schema: schema,
Entity: entity,
Model: model,
Options: options,
Data: itemMap,
Writer: w,
Tx: tx,
Context: ctx,
Handler: h,
Schema: schema,
Entity: entity,
Model: model,
Operation: "create",
Options: options,
Data: itemMap,
Writer: w,
Tx: tx,
}
if err := h.hooks.ExecuteBeforeOp(BeforeCreate, hookCtx); err != nil {
return fmt.Errorf("BeforeCreate hook failed: %w", err)
@@ -1166,16 +1177,17 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
// they must run before the existence-check select so that select is
// also subject to RLS on this connection/transaction.
hookCtx := &HookContext{
Context: ctx,
Handler: h,
Schema: schema,
Entity: entity,
Model: model,
Options: options,
ID: urlID,
Data: updates,
Writer: w,
Tx: tx,
Context: ctx,
Handler: h,
Schema: schema,
Entity: entity,
Model: model,
Operation: "update",
Options: options,
ID: urlID,
Data: updates,
Writer: w,
Tx: tx,
}
if err := h.hooks.ExecuteBeforeOp(BeforeUpdate, hookCtx); err != nil {
@@ -1387,16 +1399,17 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
// Execute BeforeUpdate hooks inside transaction
hookCtx := &HookContext{
Context: ctx,
Handler: h,
Schema: schema,
Entity: entity,
Model: model,
Options: options,
ID: itemIDStr,
Data: item,
Writer: w,
Tx: tx,
Context: ctx,
Handler: h,
Schema: schema,
Entity: entity,
Model: model,
Operation: "update",
Options: options,
ID: itemIDStr,
Data: item,
Writer: w,
Tx: tx,
}
if err := h.hooks.ExecuteBeforeOp(BeforeUpdate, hookCtx); err != nil {
@@ -1543,16 +1556,17 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
// Execute BeforeUpdate hooks inside transaction
hookCtx := &HookContext{
Context: ctx,
Handler: h,
Schema: schema,
Entity: entity,
Model: model,
Options: options,
ID: itemIDStr,
Data: itemMap,
Writer: w,
Tx: tx,
Context: ctx,
Handler: h,
Schema: schema,
Entity: entity,
Model: model,
Operation: "update",
Options: options,
ID: itemIDStr,
Data: itemMap,
Writer: w,
Tx: tx,
}
if err := h.hooks.ExecuteBeforeOp(BeforeUpdate, hookCtx); err != nil {
@@ -1648,15 +1662,16 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id
// Execute BeforeDelete hooks (covers model-rule checks before any deletion)
hookCtx := &HookContext{
Context: ctx,
Handler: h,
Schema: schema,
Entity: entity,
Model: model,
ID: id,
Data: data,
Writer: w,
Tx: h.db,
Context: ctx,
Handler: h,
Schema: schema,
Entity: entity,
Model: model,
Operation: "delete",
ID: id,
Data: data,
Writer: w,
Tx: h.db,
}
if err := h.hooks.ExecuteBeforeOp(BeforeDelete, hookCtx); err != nil {
logger.Error("BeforeDelete hook failed: %v", err)
@@ -1937,10 +1952,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("CAST(%s AS TEXT) LIKE ?", filter.Column)
condition = fmt.Sprintf("%s LIKE ?", likeColumn(filter.Column, model))
args = []interface{}{filter.Value}
case "ilike":
condition = fmt.Sprintf("CAST(%s AS TEXT) ILIKE ?", filter.Column)
condition = fmt.Sprintf("%s ILIKE ?", likeColumn(filter.Column, model))
args = []interface{}{filter.Value}
case "in":
condition, args = common.BuildInCondition(filter.Column, filter.Value)
@@ -1973,6 +1988,18 @@ func (h *Handler) buildFilterCondition(filter common.FilterOption, model interfa
return condition, args
}
// likeColumn returns the column expression to use for LIKE/ILIKE. citext
// columns are compared natively — they're already case-insensitive, and
// 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) {
return column
}
return fmt.Sprintf("CAST(%s AS TEXT)", column)
}
func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOption, model interface{}) common.SelectQuery {
// Determine which method to use based on LogicOperator
useOrLogic := strings.EqualFold(filter.LogicOperator, "OR")
@@ -2007,10 +2034,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("CAST(%s AS TEXT) LIKE ?", filter.Column)
condition = fmt.Sprintf("%s LIKE ?", likeColumn(filter.Column, model))
args = []interface{}{filter.Value}
case "ilike":
condition = fmt.Sprintf("CAST(%s AS TEXT) ILIKE ?", filter.Column)
condition = fmt.Sprintf("%s ILIKE ?", likeColumn(filter.Column, model))
args = []interface{}{filter.Value}
case "in":
condition, args = common.BuildInCondition(filter.Column, filter.Value)
+10
View File
@@ -25,12 +25,18 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList
// Hook 1: BeforeRead - Load security rules
handler.Hooks().Register(BeforeRead, func(hookCtx *HookContext) error {
secCtx := newSecurityContext(hookCtx)
if security.IsModelSecurityDisabled(secCtx) {
return nil
}
return security.LoadSecurityRules(secCtx, securityList)
})
// Hook 2: BeforeScan - Apply row-level security filters
handler.Hooks().Register(BeforeScan, func(hookCtx *HookContext) error {
secCtx := newSecurityContext(hookCtx)
if security.ShouldSkipRowSecurity(secCtx, hookCtx.Operation) {
return nil
}
return security.ApplyRowSecurity(secCtx, securityList)
})
@@ -97,6 +103,10 @@ func (s *securityContext) GetEntity() string {
return s.ctx.Entity
}
func (s *securityContext) GetOperation() string {
return s.ctx.Operation
}
func (s *securityContext) GetModel() interface{} {
return s.ctx.Model
}
+175
View File
@@ -0,0 +1,175 @@
package restheadspec
import (
"reflect"
"testing"
"github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
)
// atdetailModel mirrors the real-world model that triggered this regression:
// rid_parent is a nullable bigint foreign key, backed by spectypes.SqlInt64
// (a SqlNull[int64] alias). An eq filter on it was being rendered as
// CAST(atdetail.rid_parent AS TEXT) = '90446096', which can't use the index
// on rid_parent. Name is a citext column, which must never be cast to TEXT
// either (that would switch to case-sensitive matching and lose its index).
type atdetailModel struct {
RidParent spectypes.SqlInt64 `json:"rid_parent" bun:"rid_parent"`
Name string `json:"name" bun:"name,type:citext"`
}
func TestValidateAndAdjustFilterForColumnType_SqlNullNumeric(t *testing.T) {
h := &Handler{}
model := atdetailModel{}
filter := &common.FilterOption{Column: "rid_parent", Operator: "eq", Value: "90446096"}
info := h.ValidateAndAdjustFilterForColumnType(filter, model)
if info.NeedsCast {
t.Fatalf("expected NeedsCast=false for a numeric SqlInt64 column with a numeric value, got true")
}
if !info.IsNumericType {
t.Fatalf("expected IsNumericType=true for a SqlInt64 column")
}
if v, ok := filter.Value.(int64); !ok || v != 90446096 {
t.Fatalf("expected filter.Value to be converted to int64(90446096), got %#v", filter.Value)
}
}
func TestApplyFilter_SqlNullNumeric_NoCastKeepsIndexUsable(t *testing.T) {
h := &Handler{}
model := atdetailModel{}
filter := common.FilterOption{Column: "rid_parent", Operator: "eq", Value: "90446096"}
castInfo := h.ValidateAndAdjustFilterForColumnType(&filter, model)
q := &jsonCapQuery{}
h.applyFilter(q, filter, "public.atdetail", castInfo.NeedsCast, "AND", model)
c := q.only(t)
const want = "atdetail.rid_parent = ?"
if c.query != want {
t.Fatalf("query = %q, want %q (must not CAST a numeric column to TEXT)", c.query, want)
}
if !reflect.DeepEqual(c.args, []interface{}{int64(90446096)}) {
t.Fatalf("args = %#v", c.args)
}
}
// TestFieldFilterHeader_SqlNullNumeric_EndToEnd reproduces the exact reported
// regression: a request carrying the header
//
// x-fieldfilter-rid_parent: 90446096
//
// against a model whose rid_parent field is a nullable bigint (spectypes.SqlInt64).
// Before the fix, this parsed to a filter that got CAST(atdetail.rid_parent AS TEXT) = '90446096',
// making the query unable to use the index on rid_parent. It must now parse to
// a native "atdetail.rid_parent = ?" comparison with an int64 argument.
func TestFieldFilterHeader_SqlNullNumeric_EndToEnd(t *testing.T) {
h := NewHandler(nil, nil)
model := atdetailModel{}
req := &MockRequest{
headers: map[string]string{
"x-fieldfilter-rid_parent": "90446096",
},
queryParams: map[string]string{},
}
options := h.parseOptionsFromHeaders(req, model)
if len(options.Filters) != 1 {
t.Fatalf("expected 1 filter parsed from x-fieldfilter-rid_parent, got %d: %+v", len(options.Filters), options.Filters)
}
filter := options.Filters[0]
if filter.Column != "rid_parent" || filter.Operator != "eq" {
t.Fatalf("unexpected parsed filter: %+v", filter)
}
if filter.Value != "90446096" {
t.Fatalf("expected raw header string value before type validation, got %#v", filter.Value)
}
// This is the exact step that decided whether to CAST: ValidateAndAdjustFilterForColumnType
// used to see reflect.Struct for the SqlInt64-wrapped column and cast to TEXT.
castInfo := h.ValidateAndAdjustFilterForColumnType(&filter, model)
if castInfo.NeedsCast {
t.Fatalf("regression: numeric SqlInt64 column x-fieldfilter-rid_parent got NeedsCast=true, " +
"which renders CAST(atdetail.rid_parent AS TEXT) = '90446096' and defeats the column's index")
}
q := &jsonCapQuery{}
h.applyFilter(q, filter, "public.atdetail", castInfo.NeedsCast, filter.LogicOperator, model)
c := q.only(t)
const want = "atdetail.rid_parent = ?"
if c.query != want {
t.Fatalf("SQL condition = %q, want %q (no CAST, so the rid_parent index can still be used)", c.query, want)
}
if !reflect.DeepEqual(c.args, []interface{}{int64(90446096)}) {
t.Fatalf("args = %#v, want [int64(90446096)]", c.args)
}
}
func TestApplyFilter_Citext_NeverCastForEqOrIlike(t *testing.T) {
h := &Handler{}
model := atdetailModel{}
t.Run("eq", func(t *testing.T) {
filter := common.FilterOption{Column: "name", Operator: "eq", Value: "Acme"}
castInfo := h.ValidateAndAdjustFilterForColumnType(&filter, model)
if castInfo.NeedsCast {
t.Fatalf("citext column must never need a CAST")
}
q := &jsonCapQuery{}
h.applyFilter(q, filter, "public.atdetail", castInfo.NeedsCast, "AND", model)
if c := q.only(t); c.query != "atdetail.name = ?" {
t.Fatalf("query = %q", c.query)
}
})
t.Run("ilike", func(t *testing.T) {
filter := common.FilterOption{Column: "name", Operator: "ilike", Value: "%acme%"}
q := &jsonCapQuery{}
h.applyFilter(q, filter, "public.atdetail", false, "AND", model)
if c := q.only(t); c.query != "atdetail.name ILIKE ?" {
t.Fatalf("query = %q, want no CAST for a citext column", c.query)
}
})
}
// TestValidateAndAdjustFilterForColumnType_NumericColumn_Ilike reproduces a
// global "search all columns" request (x-searchor-contains-<col> per column,
// e.g. the X-Filter-All style OR group) landing an ILIKE filter with a
// '%...%'-wrapped numeric-looking value on a numeric column such as
// rid_parent. Before the fix, ValidateAndAdjustFilterForColumnType trimmed
// the '%' wildcards, saw a numeric string, and rewrote filter.Value to an
// int64 -- so applyFilter's CAST(col AS TEXT) ILIKE ? bound an integer
// argument instead of the wildcard string, and Postgres rejected it with
// "operator does not exist: text ~~* integer".
func TestValidateAndAdjustFilterForColumnType_NumericColumn_Ilike(t *testing.T) {
h := &Handler{}
model := atdetailModel{}
filter := &common.FilterOption{Column: "rid_parent", Operator: "ilike", Value: "%345346346%"}
info := h.ValidateAndAdjustFilterForColumnType(filter, model)
if !info.NeedsCast {
t.Fatalf("expected NeedsCast=true so the numeric column is cast to TEXT for ILIKE")
}
if filter.Value != "%345346346%" {
t.Fatalf("ILIKE must keep the wildcard-wrapped string value untouched, got %#v", filter.Value)
}
q := &jsonCapQuery{}
h.applyFilter(q, *filter, "public.atdetail", info.NeedsCast, "OR", model)
c := q.only(t)
const want = "CAST(atdetail.rid_parent AS TEXT) ILIKE ?"
if c.query != want {
t.Fatalf("query = %q, want %q", c.query, want)
}
if !reflect.DeepEqual(c.args, []interface{}{"%345346346%"}) {
t.Fatalf("args = %#v, want [\"%%345346346%%\"]", c.args)
}
}
+87 -8
View File
@@ -233,8 +233,18 @@ func (h *Handler) Handle(w common.ResponseWriter, r common.Request, params map[s
return
}
validId, _ := strconv.ParseInt(id, 10, 64)
if validId > 0 {
h.handleUpdate(ctx, w, id, nil, data, options)
updateID := id
isUpdate := validId > 0
if !isUpdate {
// No valid /:id in the URL - check whether the body itself carries
// a valid primary key value and treat this as an update if so.
if pkID, ok := h.extractPrimaryKeyFromBody(model, data); ok && pkID != "0" {
updateID = pkID
isUpdate = true
}
}
if isUpdate {
h.handleUpdate(ctx, w, updateID, nil, data, options)
} else {
h.handleCreate(ctx, w, data, options)
}
@@ -271,6 +281,49 @@ func (h *Handler) Handle(w common.ResponseWriter, r common.Request, params map[s
}
}
// extractPrimaryKeyFromBody looks for a valid primary key value inside a
// decoded (single-record) POST body, keyed by the model's primary key column
// or its JSON equivalent. It returns the string form of that value and true
// if one was found and is non-empty/non-zero; otherwise ("", false).
func (h *Handler) extractPrimaryKeyFromBody(model interface{}, data interface{}) (string, bool) {
dataMap, ok := data.(map[string]interface{})
if !ok {
// Batch payloads (slices) aren't eligible for this implicit-update detection.
return "", false
}
pkCol := reflection.GetPrimaryKeyName(model)
if pkCol == "" {
return "", false
}
val, exists := dataMap[pkCol]
if !exists {
modelType := reflection.GetPointerElement(reflect.TypeOf(model))
for jsonKey, col := range reflection.BuildJSONToDBColumnMap(modelType) {
if col == pkCol {
val, exists = dataMap[jsonKey]
break
}
}
}
if !exists || val == nil || reflection.IsEmptyValue(val) {
return "", false
}
switch v := val.(type) {
case float64:
if v <= 0 {
return "", false
}
return strconv.FormatInt(int64(v), 10), true
case string:
return v, true
default:
return fmt.Sprintf("%v", v), true
}
}
// HandleGet processes GET requests for metadata
func (h *Handler) HandleGet(w common.ResponseWriter, r common.Request, params map[string]string) {
// Capture panics and return error response
@@ -379,6 +432,7 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
Entity: entity,
TableName: tableName,
Model: model,
Operation: "read",
Options: options,
ID: id,
Writer: w,
@@ -476,6 +530,10 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
// JSON sub-field selection (data->>'x', data.x, data#>>'{a,b}'):
// emit a parameterised expression aliased to a stable name.
if expr, jargs, alias, ok := common.ResolveJSONColumnExpr(model, selectAlias, col); ok {
if !reflection.HasColumn(model, alias) {
logger.Warn("Skipping JSON select column %q: model has no scan target for alias %q", col, alias)
continue
}
query = query.ColumnExpr(expr+" AS "+common.QuoteIdent(alias), jargs...)
continue
}
@@ -1236,6 +1294,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
Entity: entity,
TableName: tableName,
Model: model,
Operation: "create",
Options: options,
Data: data,
Writer: w,
@@ -1335,6 +1394,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
Entity: entity,
TableName: tableName,
Model: model,
Operation: "create",
Options: options,
Data: modelValue,
Writer: w,
@@ -1489,6 +1549,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id
TableName: tableName,
Tx: tx,
Model: model,
Operation: "update",
Options: options,
ID: id,
Data: dataMap,
@@ -1686,6 +1747,7 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id
Entity: entity,
TableName: tableName,
Model: model,
Operation: "delete",
ID: itemID,
Writer: w,
Tx: tx,
@@ -1760,6 +1822,7 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id
Entity: entity,
TableName: tableName,
Model: model,
Operation: "delete",
ID: itemIDStr,
Writer: w,
Tx: tx,
@@ -1818,6 +1881,7 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id
Entity: entity,
TableName: tableName,
Model: model,
Operation: "delete",
ID: itemIDStr,
Writer: w,
Tx: tx,
@@ -1902,6 +1966,7 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id
Entity: entity,
TableName: tableName,
Model: model,
Operation: "delete",
ID: id,
Writer: w,
Tx: h.db,
@@ -2325,6 +2390,13 @@ func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOpti
qualifiedColumn = fmt.Sprintf("CAST(%s AS TEXT)", rawQualifiedColumn)
}
// citext columns already compare case-insensitively; casting to TEXT for
// LIKE/ILIKE would switch to case-sensitive matching and defeat a citext index.
likeColumn := rawQualifiedColumn
if !reflection.IsCitextColumn(model, filter.Column) {
likeColumn = fmt.Sprintf("CAST(%s AS TEXT)", rawQualifiedColumn)
}
switch strings.ToLower(filter.Operator) {
case "eq", "equals":
return applyWhere(fmt.Sprintf("%s = ?", qualifiedColumn), filter.Value)
@@ -2339,11 +2411,14 @@ func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOpti
case "lte", "less_than_equals", "le":
return applyWhere(fmt.Sprintf("%s <= ?", qualifiedColumn), filter.Value)
case "like":
// Always cast to TEXT for LIKE/ILIKE to support date/time/timestamp columns
return applyWhere(fmt.Sprintf("CAST(%s AS TEXT) LIKE ?", rawQualifiedColumn), filter.Value)
// Cast to TEXT for LIKE to support date/time/timestamp columns; citext
// columns are compared natively (see likeColumn above).
return applyWhere(fmt.Sprintf("%s LIKE ?", likeColumn), filter.Value)
case "ilike":
// Always cast to TEXT for LIKE/ILIKE to support date/time/timestamp columns
return applyWhere(fmt.Sprintf("CAST(%s AS TEXT) ILIKE ?", rawQualifiedColumn), filter.Value)
// Cast to TEXT for ILIKE to support date/time/timestamp columns; citext
// columns are compared natively (see likeColumn above) since citext is
// already case-insensitive.
return applyWhere(fmt.Sprintf("%s ILIKE ?", likeColumn), filter.Value)
case "in":
cond, inArgs := common.BuildInCondition(qualifiedColumn, filter.Value)
if cond == "" {
@@ -2421,8 +2496,12 @@ func (h *Handler) applyOrFilterGroup(query common.SelectQuery, filters []*common
op := strings.ToLower(filter.Operator)
if op == "like" || op == "ilike" {
// Always cast to TEXT for LIKE/ILIKE to support date/time/timestamp columns
qualifiedColumn = fmt.Sprintf("CAST(%s AS TEXT)", rawQualifiedColumn)
// Cast to TEXT for LIKE/ILIKE to support date/time/timestamp columns.
// citext columns are left native: they're already case-insensitive and
// casting would defeat a citext index.
if !reflection.IsCitextColumn(model, filter.Column) {
qualifiedColumn = fmt.Sprintf("CAST(%s AS TEXT)", rawQualifiedColumn)
}
} else if castInfo[i].NeedsCast {
// Apply casting to text if needed for non-numeric columns or non-numeric values
qualifiedColumn = fmt.Sprintf("CAST(%s AS TEXT)", rawQualifiedColumn)
+18
View File
@@ -1466,6 +1466,12 @@ func (h *Handler) ValidateAndAdjustFilterForColumnType(filter *common.FilterOpti
return ColumnCastInfo{NeedsCast: false, IsNumericType: false}
}
// Never cast citext columns to TEXT: CAST(col AS TEXT) swaps in case-sensitive
// comparison semantics and prevents PostgreSQL from using a citext index.
if reflection.IsCitextColumn(model, filter.Column) {
return ColumnCastInfo{NeedsCast: false, IsNumericType: false}
}
colType := reflection.GetColumnTypeFromModel(model, filter.Column)
if colType == reflect.Invalid {
// Column not found in model, no casting needed
@@ -1473,6 +1479,18 @@ func (h *Handler) ValidateAndAdjustFilterForColumnType(filter *common.FilterOpti
return ColumnCastInfo{NeedsCast: false, IsNumericType: false}
}
// LIKE/ILIKE always compare against text, wildcards and all. Never coerce
// the value to the column's native numeric/bool/time type here: doing so
// strips the '%' wildcards and hands the driver a non-string argument,
// which fails with "operator does not exist: text ~~* integer" once the
// column is cast to TEXT below.
if op := strings.ToLower(filter.Operator); op == "like" || op == "ilike" {
if reflection.IsStringType(colType) {
return ColumnCastInfo{NeedsCast: false, IsNumericType: false}
}
return ColumnCastInfo{NeedsCast: true, IsNumericType: reflection.IsNumericType(colType)}
}
// Check if the input value is numeric
valueIsNumeric := false
if strVal, ok := filter.Value.(string); ok {
+59 -17
View File
@@ -232,9 +232,37 @@ func LoadSecurityRules(secCtx SecurityContext, securityList *SecurityList) error
// ApplyRowSecurity is a public wrapper for applyRowSecurity that accepts a SecurityContext
// This allows other packages to apply row-level security using the generic interface
func ApplyRowSecurity(secCtx SecurityContext, securityList *SecurityList) error {
// Spec adapters that expose the dispatched operation can enforce the same
// model-rule bypass even when ApplyRowSecurity is called directly.
if operationCtx, ok := secCtx.(interface{ GetOperation() string }); ok &&
ShouldSkipRowSecurity(secCtx, operationCtx.GetOperation()) {
return nil
}
return applyRowSecurity(secCtx, securityList)
}
// ShouldSkipRowSecurity reports whether row-security enforcement should be
// skipped for the operation. It uses the same model-rule resolution as
// CheckModelAuthAllowed so the model registry remains the single source of
// truth for security behavior.
func ShouldSkipRowSecurity(secCtx SecurityContext, operation string) bool {
rules, ok := resolveModelRules(secCtx)
if !ok {
return false
}
return rules.SecurityDisabled || (operation == "read" && rules.CanPublicRead)
}
// IsModelSecurityDisabled reports whether all model-level security processing
// is disabled for the model. This is distinct from ShouldSkipRowSecurity:
// CanPublicRead skips row filtering for reads but must still allow other read
// security, such as column masking, to be loaded.
func IsModelSecurityDisabled(secCtx SecurityContext) bool {
rules, ok := resolveModelRules(secCtx)
return ok && rules.SecurityDisabled
}
// ApplyColumnSecurity is a public wrapper for applyColumnSecurity that accepts a SecurityContext
// This allows other packages to apply column-level security using the generic interface
func ApplyColumnSecurity(secCtx SecurityContext, securityList *SecurityList) error {
@@ -303,25 +331,14 @@ func checkModelDeleteAllowed(secCtx SecurityContext) error {
// 7. Guest (UserID == 0) → return "authentication required".
// 8. Authenticated user → allow (operation-specific checks remain in BeforeUpdate/BeforeDelete).
func CheckModelAuthAllowed(secCtx SecurityContext, operation string) error {
rules, ok := GetModelRulesFromContext(secCtx.GetContext())
rules, ok := resolveModelRules(secCtx)
if !ok {
schema := secCtx.GetSchema()
entity := secCtx.GetEntity()
var err error
if schema != "" {
rules, err = modelregistry.GetModelRulesByName(fmt.Sprintf("%s.%s", schema, entity))
}
if err != nil || schema == "" {
rules, err = modelregistry.GetModelRulesByName(entity)
}
if err != nil {
// Model not registered - fall through to auth check
userID, _ := secCtx.GetUserID()
if userID == 0 {
return fmt.Errorf("authentication required")
}
return nil
// Model not registered - fall through to auth check
userID, _ := secCtx.GetUserID()
if userID == 0 {
return fmt.Errorf("authentication required")
}
return nil
}
if rules.SecurityDisabled {
@@ -347,6 +364,31 @@ func CheckModelAuthAllowed(secCtx SecurityContext, operation string) error {
return nil
}
// resolveModelRules returns model rules from the request context first, then
// falls back to the schema-qualified and unqualified registry names.
func resolveModelRules(secCtx SecurityContext) (modelregistry.ModelRules, bool) {
if rules, ok := GetModelRulesFromContext(secCtx.GetContext()); ok {
return rules, true
}
schema := secCtx.GetSchema()
entity := secCtx.GetEntity()
var err error
if schema != "" {
var rules modelregistry.ModelRules
rules, err = modelregistry.GetModelRulesByName(fmt.Sprintf("%s.%s", schema, entity))
if err == nil {
return rules, true
}
}
rules, err := modelregistry.GetModelRulesByName(entity)
if err != nil {
return modelregistry.ModelRules{}, false
}
return rules, true
}
// CheckModelUpdateAllowed is the public wrapper for checkModelUpdateAllowed.
func CheckModelUpdateAllowed(secCtx SecurityContext) error {
return checkModelUpdateAllowed(secCtx)
+249
View File
@@ -0,0 +1,249 @@
// Package quickproxy provides a small reverse-proxy layer that tries a set
// of configured upstream targets first, and falls back to a caller-supplied
// http.Handler (typically static file serving) when the upstream is
// unreachable or returns 404.
package quickproxy
import (
"bytes"
"errors"
"fmt"
"io"
"net"
"net/http"
"net/http/httputil"
"net/url"
"sort"
"strings"
"time"
)
// Rule maps a URL path prefix to an upstream target.
// A Rule with URLPrefix "/" acts as a catch-all passthrough.
type Rule struct {
// URLPrefix is the URL path prefix this rule matches. Must start with "/".
URLPrefix string
// Target is the upstream base URL, e.g. "http://localhost:3000".
// The incoming request path and query are forwarded unchanged; only the
// scheme and host are rewritten to Target's.
Target string
// Exclude is a list of URL path prefixes that this rule should not
// proxy, even though they fall under URLPrefix. Each entry is a full
// path from root and must itself start with URLPrefix (e.g. rule
// URLPrefix "/api" excluding a subpath must use "/api/health", not
// "/health"). A request matching an Exclude prefix is treated as if
// this rule didn't match at all: matching continues against any other
// configured rule, falling back if none match. This is typically used
// to carve out paths (e.g. "/health") from a catch-all "/" rule so
// they're served by the fallback handler instead of being proxied.
Exclude []string
}
// DefaultTimeout is the dial and response-header timeout applied to
// upstream requests when no WithTimeout option is given. It does not limit
// response body streaming.
const DefaultTimeout = 10 * time.Second
// Option configures a Service.
type Option func(*options)
type options struct {
timeout time.Duration
}
// WithTimeout sets the dial and response-header timeout used when
// connecting to upstream targets. It does not limit response body
// streaming, so it won't interrupt long-lived downloads or SSE/WebSocket
// connections once established.
func WithTimeout(d time.Duration) Option {
return func(o *options) { o.timeout = d }
}
// compiledRule pairs a Rule with its ready-to-use reverse proxy.
type compiledRule struct {
prefix string
excludes []string
proxy *httputil.ReverseProxy
}
// excluded reports whether path falls under one of the rule's Exclude prefixes.
func (r *compiledRule) excluded(path string) bool {
for _, ex := range r.excludes {
if strings.HasPrefix(path, ex) {
return true
}
}
return false
}
// Service holds a compiled set of proxy rules and performs longest-prefix
// matching against them. A Service is safe for concurrent use once
// returned from NewService; Handler must be called once per Service to
// wire up the fallback handler before the returned http.Handler is served.
type Service struct {
rules []compiledRule // sorted by descending prefix length
}
// errUpstreamNotFound is a sentinel error returned from ModifyResponse to
// make ReverseProxy invoke ErrorHandler (our fallback path) instead of
// writing the upstream's 404 to the client. Nothing has been written to
// the ResponseWriter yet when this happens.
var errUpstreamNotFound = errors.New("quickproxy: upstream returned 404")
// NewService compiles the given rules into a Service. Rules are matched by
// longest URLPrefix, so a catch-all "/" rule can coexist with more specific
// rules such as "/api".
func NewService(rules []Rule, opts ...Option) (*Service, error) {
if len(rules) == 0 {
return nil, fmt.Errorf("quickproxy: no rules configured")
}
cfg := options{timeout: DefaultTimeout}
for _, opt := range opts {
opt(&cfg)
}
seen := make(map[string]bool, len(rules))
compiled := make([]compiledRule, 0, len(rules))
for _, r := range rules {
if !strings.HasPrefix(r.URLPrefix, "/") {
return nil, fmt.Errorf("quickproxy: rule prefix %q must start with /", r.URLPrefix)
}
if seen[r.URLPrefix] {
return nil, fmt.Errorf("quickproxy: duplicate rule prefix %q", r.URLPrefix)
}
seen[r.URLPrefix] = true
target, err := url.Parse(r.Target)
if err != nil || target.Scheme == "" || target.Host == "" {
return nil, fmt.Errorf("quickproxy: invalid target %q for prefix %q", r.Target, r.URLPrefix)
}
for _, ex := range r.Exclude {
if !strings.HasPrefix(ex, "/") {
return nil, fmt.Errorf("quickproxy: exclude prefix %q for rule %q must start with /", ex, r.URLPrefix)
}
if !strings.HasPrefix(ex, r.URLPrefix) {
return nil, fmt.Errorf("quickproxy: exclude prefix %q for rule %q must itself start with the rule's URLPrefix", ex, r.URLPrefix)
}
}
compiled = append(compiled, compiledRule{
prefix: r.URLPrefix,
excludes: r.Exclude,
proxy: newReverseProxy(target, cfg.timeout),
})
}
// Longest prefix first, so the first match in Handler is always the
// most specific one.
sort.Slice(compiled, func(i, j int) bool {
return len(compiled[i].prefix) > len(compiled[j].prefix)
})
return &Service{rules: compiled}, nil
}
func newReverseProxy(target *url.URL, timeout time.Duration) *httputil.ReverseProxy {
transport := &http.Transport{
DialContext: (&net.Dialer{
Timeout: timeout,
}).DialContext,
ResponseHeaderTimeout: timeout,
}
return &httputil.ReverseProxy{
Transport: transport,
Director: func(req *http.Request) {
originalHost := req.Host
req.URL.Scheme = target.Scheme
req.URL.Host = target.Host
req.Host = target.Host
if originalHost != "" {
req.Header.Set("X-Forwarded-Host", originalHost)
}
},
ModifyResponse: func(resp *http.Response) error {
if resp.StatusCode == http.StatusNotFound {
return errUpstreamNotFound
}
return nil
},
}
}
// Handler returns an http.Handler that tries the configured proxy rules
// first (longest-prefix match), and calls fallback when no rule matches,
// the upstream is unreachable, or the upstream returns 404. Any other
// upstream response (2xx, other 4xx, 5xx) is streamed through to the
// client unchanged.
//
// Handler wires up ErrorHandler on the Service's compiled rules, so it
// should be called once per Service, before the returned http.Handler
// starts serving requests.
func (s *Service) Handler(fallback http.Handler) http.Handler {
if fallback == nil {
fallback = http.NotFoundHandler()
}
for i := range s.rules {
s.rules[i].proxy.ErrorHandler = func(w http.ResponseWriter, r *http.Request, _ error) {
// ReverseProxy consumes and closes r.Body while attempting the
// upstream request, even when that attempt fails (per the
// http.RoundTripper contract). Restore a fresh copy from
// r.GetBody, set below, before handing the request to fallback.
if r.GetBody != nil {
if body, err := r.GetBody(); err == nil {
r.Body = body
}
}
fallback.ServeHTTP(w, r)
}
}
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
rule := s.match(r.URL.Path)
if rule == nil {
fallback.ServeHTTP(w, r)
return
}
// Buffer the body so it can be replayed to fallback if the upstream
// attempt fails; see ErrorHandler above.
if r.Body != nil && r.Body != http.NoBody {
bodyBytes, err := io.ReadAll(r.Body)
r.Body.Close()
if err != nil {
http.Error(w, "failed to read request body", http.StatusInternalServerError)
return
}
r.Body = io.NopCloser(bytes.NewReader(bodyBytes))
r.GetBody = func() (io.ReadCloser, error) {
return io.NopCloser(bytes.NewReader(bodyBytes)), nil
}
}
rule.proxy.ServeHTTP(w, r)
})
}
// match returns the longest-prefix rule matching path, or nil if none match.
// A rule whose Exclude covers path is skipped, and matching continues
// against the next-longest-prefix rule.
func (s *Service) match(path string) *compiledRule {
for i := range s.rules {
if !strings.HasPrefix(path, s.rules[i].prefix) {
continue
}
if s.rules[i].excluded(path) {
continue
}
return &s.rules[i]
}
return nil
}
+392
View File
@@ -0,0 +1,392 @@
package quickproxy
import (
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
)
func TestNewService_Validation(t *testing.T) {
tests := []struct {
name string
rules []Rule
wantErr bool
}{
{"no rules", nil, true},
{"empty rules", []Rule{}, true},
{"bad prefix", []Rule{{URLPrefix: "api", Target: "http://localhost:1"}}, true},
{"bad target", []Rule{{URLPrefix: "/api", Target: "not-a-url"}}, true},
{"missing host", []Rule{{URLPrefix: "/api", Target: "http://"}}, true},
{"duplicate prefix", []Rule{
{URLPrefix: "/api", Target: "http://localhost:1"},
{URLPrefix: "/api", Target: "http://localhost:2"},
}, true},
{"bad exclude prefix", []Rule{
{URLPrefix: "/", Target: "http://localhost:1", Exclude: []string{"health"}},
}, true},
{"exclude outside rule's URLPrefix", []Rule{
{URLPrefix: "/api", Target: "http://localhost:1", Exclude: []string{"/health"}},
}, true},
{"valid", []Rule{{URLPrefix: "/api", Target: "http://localhost:1"}}, false},
{"valid with exclude", []Rule{
{URLPrefix: "/", Target: "http://localhost:1", Exclude: []string{"/health"}},
}, false},
{"valid with nested exclude", []Rule{
{URLPrefix: "/api", Target: "http://localhost:1", Exclude: []string{"/api/health"}},
}, false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
_, err := NewService(tt.rules)
if (err != nil) != tt.wantErr {
t.Fatalf("NewService() error = %v, wantErr %v", err, tt.wantErr)
}
})
}
}
func fallbackHandler(body string) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte(body))
})
}
func TestHandler_ProxiesSuccessResponse(t *testing.T) {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte("upstream:" + r.URL.Path))
}))
defer upstream.Close()
svc, err := NewService([]Rule{{URLPrefix: "/api", Target: upstream.URL}})
if err != nil {
t.Fatalf("NewService: %v", err)
}
handler := svc.Handler(fallbackHandler("fallback"))
req := httptest.NewRequest(http.MethodGet, "/api/widgets", nil)
rr := httptest.NewRecorder()
handler.ServeHTTP(rr, req)
if rr.Code != http.StatusOK {
t.Fatalf("status = %d, want 200", rr.Code)
}
if got := rr.Body.String(); got != "upstream:/api/widgets" {
t.Fatalf("body = %q", got)
}
}
func TestHandler_404FallsBack(t *testing.T) {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusNotFound)
_, _ = w.Write([]byte("upstream not found"))
}))
defer upstream.Close()
svc, err := NewService([]Rule{{URLPrefix: "/", Target: upstream.URL}})
if err != nil {
t.Fatalf("NewService: %v", err)
}
handler := svc.Handler(fallbackHandler("fallback-content"))
req := httptest.NewRequest(http.MethodGet, "/missing.html", nil)
rr := httptest.NewRecorder()
handler.ServeHTTP(rr, req)
if rr.Code != http.StatusOK {
t.Fatalf("status = %d, want 200", rr.Code)
}
if got := rr.Body.String(); got != "fallback-content" {
t.Fatalf("body = %q, want fallback-content", got)
}
}
func TestHandler_UnreachableUpstreamFallsBack(t *testing.T) {
// A closed listener address: nothing is listening, so dialing fails.
unreachable := "http://127.0.0.1:1"
svc, err := NewService([]Rule{{URLPrefix: "/", Target: unreachable}}, WithTimeout(500*time.Millisecond))
if err != nil {
t.Fatalf("NewService: %v", err)
}
handler := svc.Handler(fallbackHandler("fallback-content"))
req := httptest.NewRequest(http.MethodGet, "/anything", nil)
rr := httptest.NewRecorder()
handler.ServeHTTP(rr, req)
if rr.Code != http.StatusOK {
t.Fatalf("status = %d, want 200", rr.Code)
}
if got := rr.Body.String(); got != "fallback-content" {
t.Fatalf("body = %q, want fallback-content", got)
}
}
func TestHandler_UnreachableUpstreamFallsBackWithBody(t *testing.T) {
// A closed listener address: nothing is listening, so dialing fails and
// ReverseProxy invokes ErrorHandler. The fallback handler must still see
// the original request body, even though ReverseProxy consumed and
// closed it while attempting (and failing) the upstream request.
unreachable := "http://127.0.0.1:1"
svc, err := NewService([]Rule{{URLPrefix: "/", Target: unreachable}}, WithTimeout(500*time.Millisecond))
if err != nil {
t.Fatalf("NewService: %v", err)
}
echoBody := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
body, err := io.ReadAll(r.Body)
if err != nil {
t.Fatalf("fallback reading body: %v", err)
}
w.WriteHeader(http.StatusOK)
_, _ = w.Write(body)
})
handler := svc.Handler(echoBody)
req := httptest.NewRequest(http.MethodPost, "/submit", strings.NewReader("payload=1"))
rr := httptest.NewRecorder()
handler.ServeHTTP(rr, req)
if rr.Code != http.StatusOK {
t.Fatalf("status = %d, want 200", rr.Code)
}
if got := rr.Body.String(); got != "payload=1" {
t.Fatalf("body = %q, want payload=1", got)
}
}
func TestHandler_404FallsBackWithBody(t *testing.T) {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusNotFound)
}))
defer upstream.Close()
svc, err := NewService([]Rule{{URLPrefix: "/", Target: upstream.URL}})
if err != nil {
t.Fatalf("NewService: %v", err)
}
echoBody := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
body, err := io.ReadAll(r.Body)
if err != nil {
t.Fatalf("fallback reading body: %v", err)
}
w.WriteHeader(http.StatusOK)
_, _ = w.Write(body)
})
handler := svc.Handler(echoBody)
req := httptest.NewRequest(http.MethodPut, "/missing", strings.NewReader("payload=2"))
rr := httptest.NewRecorder()
handler.ServeHTTP(rr, req)
if rr.Code != http.StatusOK {
t.Fatalf("status = %d, want 200", rr.Code)
}
if got := rr.Body.String(); got != "payload=2" {
t.Fatalf("body = %q, want payload=2", got)
}
}
func TestHandler_NonNotFoundErrorsPassThrough(t *testing.T) {
codes := []int{http.StatusOK, http.StatusForbidden, http.StatusBadRequest, http.StatusInternalServerError}
for _, code := range codes {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(code)
_, _ = w.Write([]byte("upstream response"))
}))
svc, err := NewService([]Rule{{URLPrefix: "/", Target: upstream.URL}})
if err != nil {
upstream.Close()
t.Fatalf("NewService: %v", err)
}
handler := svc.Handler(fallbackHandler("fallback-content"))
req := httptest.NewRequest(http.MethodGet, "/x", nil)
rr := httptest.NewRecorder()
handler.ServeHTTP(rr, req)
if rr.Code != code {
t.Errorf("status for upstream code %d = %d, want %d", code, rr.Code, code)
}
if got := rr.Body.String(); got != "upstream response" {
t.Errorf("body for upstream code %d = %q, want passthrough", code, got)
}
upstream.Close()
}
}
func TestHandler_LongestPrefixMatch(t *testing.T) {
specific := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
_, _ = w.Write([]byte("specific"))
}))
defer specific.Close()
general := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
_, _ = w.Write([]byte("general"))
}))
defer general.Close()
svc, err := NewService([]Rule{
{URLPrefix: "/", Target: general.URL},
{URLPrefix: "/api/v1", Target: specific.URL},
})
if err != nil {
t.Fatalf("NewService: %v", err)
}
handler := svc.Handler(fallbackHandler("fallback"))
for path, want := range map[string]string{
"/api/v1/thing": "specific",
"/api/other": "general",
"/anything": "general",
} {
req := httptest.NewRequest(http.MethodGet, path, nil)
rr := httptest.NewRecorder()
handler.ServeHTTP(rr, req)
if got := rr.Body.String(); got != want {
t.Errorf("path %s: body = %q, want %q", path, got, want)
}
}
}
func TestHandler_ExcludeFallsBackToFallback(t *testing.T) {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
_, _ = w.Write([]byte("upstream:" + r.URL.Path))
}))
defer upstream.Close()
svc, err := NewService([]Rule{
{URLPrefix: "/", Target: upstream.URL, Exclude: []string{"/health"}},
})
if err != nil {
t.Fatalf("NewService: %v", err)
}
handler := svc.Handler(fallbackHandler("fallback-content"))
req := httptest.NewRequest(http.MethodGet, "/health", nil)
rr := httptest.NewRecorder()
handler.ServeHTTP(rr, req)
if got := rr.Body.String(); got != "fallback-content" {
t.Fatalf("body = %q, want fallback-content", got)
}
req = httptest.NewRequest(http.MethodGet, "/health/live", nil)
rr = httptest.NewRecorder()
handler.ServeHTTP(rr, req)
if got := rr.Body.String(); got != "fallback-content" {
t.Fatalf("body = %q, want fallback-content", got)
}
req = httptest.NewRequest(http.MethodGet, "/other", nil)
rr = httptest.NewRecorder()
handler.ServeHTTP(rr, req)
if got := rr.Body.String(); got != "upstream:/other" {
t.Fatalf("body = %q, want upstream:/other", got)
}
}
func TestHandler_ExcludeFallsThroughToNextRule(t *testing.T) {
specific := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
_, _ = w.Write([]byte("specific"))
}))
defer specific.Close()
general := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
_, _ = w.Write([]byte("general"))
}))
defer general.Close()
svc, err := NewService([]Rule{
{URLPrefix: "/api", Target: general.URL},
{URLPrefix: "/api/v1", Target: specific.URL, Exclude: []string{"/api/v1/health"}},
})
if err != nil {
t.Fatalf("NewService: %v", err)
}
handler := svc.Handler(fallbackHandler("fallback"))
for path, want := range map[string]string{
"/api/v1/thing": "specific",
"/api/v1/health": "general",
} {
req := httptest.NewRequest(http.MethodGet, path, nil)
rr := httptest.NewRecorder()
handler.ServeHTTP(rr, req)
if got := rr.Body.String(); got != want {
t.Errorf("path %s: body = %q, want %q", path, got, want)
}
}
}
func TestHandler_NoMatchFallsBack(t *testing.T) {
svc, err := NewService([]Rule{{URLPrefix: "/api", Target: "http://127.0.0.1:1"}})
if err != nil {
t.Fatalf("NewService: %v", err)
}
handler := svc.Handler(fallbackHandler("fallback-content"))
req := httptest.NewRequest(http.MethodGet, "/other", nil)
rr := httptest.NewRecorder()
handler.ServeHTTP(rr, req)
if rr.Code != http.StatusOK {
t.Fatalf("status = %d, want 200", rr.Code)
}
if got := rr.Body.String(); got != "fallback-content" {
t.Fatalf("body = %q, want fallback-content", got)
}
}
func TestHandler_AllMethodsProxied(t *testing.T) {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
body, _ := io.ReadAll(r.Body)
w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte(r.Method + ":" + string(body)))
}))
defer upstream.Close()
svc, err := NewService([]Rule{{URLPrefix: "/api", Target: upstream.URL}})
if err != nil {
t.Fatalf("NewService: %v", err)
}
handler := svc.Handler(fallbackHandler("fallback"))
methods := []string{http.MethodGet, http.MethodPost, http.MethodPut, http.MethodPatch, http.MethodDelete}
for _, method := range methods {
req := httptest.NewRequest(method, "/api/widgets", nil)
rr := httptest.NewRecorder()
handler.ServeHTTP(rr, req)
want := method + ":"
if got := rr.Body.String(); got != want {
t.Errorf("method %s: body = %q, want %q", method, got, want)
}
}
}
+177
View File
@@ -0,0 +1,177 @@
package spectypes
import (
"database/sql/driver"
"encoding/json"
"fmt"
"strings"
)
// CIString is a string that stores, scans, and returns its value exactly as
// given (no case normalization), but compares case-insensitively via Equal
// and EqualString. Use it as a bun model field type for columns (e.g.
// citext, or codes matched case-insensitively) where you want Go-side
// case-insensitive comparisons without forcing the stored/returned value to
// a particular case.
type CIString string
// Value implements driver.Valuer. The value is passed through unchanged.
func (s CIString) Value() (driver.Value, error) {
return string(s), nil
}
// Scan implements sql.Scanner. The value is stored unchanged.
func (s *CIString) Scan(value any) error {
switch v := value.(type) {
case string:
*s = CIString(v)
case []byte:
*s = CIString(v)
case nil:
*s = ""
default:
return fmt.Errorf("cannot scan %T into CIString", value)
}
return nil
}
// String implements fmt.Stringer.
func (s CIString) String() string { return string(s) }
// Equal reports whether s and other are equal, ignoring case.
func (s CIString) Equal(other CIString) bool {
return strings.EqualFold(string(s), string(other))
}
// EqualString reports whether s equals other, ignoring case.
func (s CIString) EqualString(other string) bool {
return strings.EqualFold(string(s), other)
}
// Compare returns -1, 0, or +1 if s is less than, equal to, or greater than
// other, ignoring case. Useful with slices.SortFunc or similar.
func (s CIString) Compare(other CIString) int {
return strings.Compare(strings.ToLower(string(s)), strings.ToLower(string(other)))
}
// Less reports whether s sorts before other, ignoring case. Suitable for
// sort.Slice or slices.SortFunc comparisons.
func (s CIString) Less(other CIString) bool {
return s.Compare(other) < 0
}
// LCString is a string that always stores, scans, and returns as lowercase.
// Use it as a bun model field type for columns that must be normalized to
// lowercase (e.g. codes, slugs, emails) rather than merely compared
// case-insensitively; see CIString if the original case must be preserved.
type LCString string
// Value implements driver.Valuer, always lowercase.
func (s LCString) Value() (driver.Value, error) {
return strings.ToLower(string(s)), nil
}
// Scan implements sql.Scanner, always lowercase.
func (s *LCString) Scan(value any) error {
switch v := value.(type) {
case string:
*s = LCString(strings.ToLower(v))
case []byte:
*s = LCString(strings.ToLower(string(v)))
case nil:
*s = ""
default:
return fmt.Errorf("cannot scan %T into LCString", value)
}
return nil
}
// String implements fmt.Stringer, always lowercase.
func (s LCString) String() string { return strings.ToLower(string(s)) }
// Equal reports whether s and other are equal (case-insensitively, since
// both normalize to lowercase).
func (s LCString) Equal(other LCString) bool {
return s.String() == other.String()
}
// EqualString reports whether s equals other, ignoring case.
func (s LCString) EqualString(other string) bool {
return s.String() == strings.ToLower(other)
}
// MarshalJSON implements json.Marshaler, always lowercase. Needed because
// encoding/json marshals a bare string-kind type as-is and does not call
// Value/String, so a value constructed directly (not scanned from the DB)
// would otherwise serialize with its original case.
func (s LCString) MarshalJSON() ([]byte, error) {
return json.Marshal(strings.ToLower(string(s)))
}
// UnmarshalJSON implements json.Unmarshaler, always lowercase.
func (s *LCString) UnmarshalJSON(b []byte) error {
var str string
if err := json.Unmarshal(b, &str); err != nil {
return err
}
*s = LCString(strings.ToLower(str))
return nil
}
// UCString is a string that always stores, scans, and returns as uppercase.
// Use it as a bun model field type for columns that must be normalized to
// uppercase (e.g. table prefix codes) rather than merely compared
// case-insensitively; see CIString if the original case must be preserved.
type UCString string
// Value implements driver.Valuer, always uppercase.
func (s UCString) Value() (driver.Value, error) {
return strings.ToUpper(string(s)), nil
}
// Scan implements sql.Scanner, always uppercase.
func (s *UCString) Scan(value any) error {
switch v := value.(type) {
case string:
*s = UCString(strings.ToUpper(v))
case []byte:
*s = UCString(strings.ToUpper(string(v)))
case nil:
*s = ""
default:
return fmt.Errorf("cannot scan %T into UCString", value)
}
return nil
}
// String implements fmt.Stringer, always uppercase.
func (s UCString) String() string { return strings.ToUpper(string(s)) }
// Equal reports whether s and other are equal (case-insensitively, since
// both normalize to uppercase).
func (s UCString) Equal(other UCString) bool {
return s.String() == other.String()
}
// EqualString reports whether s equals other, ignoring case.
func (s UCString) EqualString(other string) bool {
return s.String() == strings.ToUpper(other)
}
// MarshalJSON implements json.Marshaler, always uppercase. Needed because
// encoding/json marshals a bare string-kind type as-is and does not call
// Value/String, so a value constructed directly (not scanned from the DB)
// would otherwise serialize with its original case.
func (s UCString) MarshalJSON() ([]byte, error) {
return json.Marshal(strings.ToUpper(string(s)))
}
// UnmarshalJSON implements json.Unmarshaler, always uppercase.
func (s *UCString) UnmarshalJSON(b []byte) error {
var str string
if err := json.Unmarshal(b, &str); err != nil {
return err
}
*s = UCString(strings.ToUpper(str))
return nil
}
+379
View File
@@ -0,0 +1,379 @@
package spectypes
import (
"encoding/json"
"sort"
"testing"
)
func TestCIString_Scan(t *testing.T) {
tests := []struct {
name string
input interface{}
expected CIString
}{
{name: "plain string", input: "MixedCase", expected: "MixedCase"},
{name: "bytes as string", input: []byte("FromBytes"), expected: "FromBytes"},
{name: "nil value", input: nil, expected: ""},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
var s CIString
if err := s.Scan(tt.input); err != nil {
t.Fatalf("Scan failed: %v", err)
}
if s != tt.expected {
t.Errorf("expected %q, got %q", tt.expected, s)
}
})
}
}
func TestCIString_Scan_InvalidType(t *testing.T) {
var s CIString
if err := s.Scan(123); err == nil {
t.Fatal("expected error scanning int into CIString, got nil")
}
}
func TestCIString_Value(t *testing.T) {
s := CIString("MixedCase")
v, err := s.Value()
if err != nil {
t.Fatalf("Value failed: %v", err)
}
if v != "MixedCase" {
t.Errorf("expected %q, got %q (case must be preserved)", "MixedCase", v)
}
}
func TestCIString_String(t *testing.T) {
s := CIString("MixedCase")
if s.String() != "MixedCase" {
t.Errorf("expected %q, got %q", "MixedCase", s.String())
}
}
func TestCIString_Equal(t *testing.T) {
tests := []struct {
name string
a, b CIString
expected bool
}{
{name: "same case", a: "ABC", b: "ABC", expected: true},
{name: "different case", a: "ABC", b: "abc", expected: true},
{name: "mixed case", a: "AbC", b: "aBc", expected: true},
{name: "not equal", a: "ABC", b: "XYZ", expected: false},
{name: "both empty", a: "", b: "", expected: true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := tt.a.Equal(tt.b); got != tt.expected {
t.Errorf("Equal(%q, %q) = %v, want %v", tt.a, tt.b, got, tt.expected)
}
})
}
}
func TestCIString_EqualString(t *testing.T) {
s := CIString("ABC")
if !s.EqualString("abc") {
t.Error("expected EqualString to match case-insensitively")
}
if s.EqualString("xyz") {
t.Error("expected EqualString to not match different strings")
}
}
func TestCIString_Compare(t *testing.T) {
tests := []struct {
name string
a, b CIString
expected int
}{
{name: "equal same case", a: "abc", b: "abc", expected: 0},
{name: "equal different case", a: "ABC", b: "abc", expected: 0},
{name: "less", a: "abc", b: "xyz", expected: -1},
{name: "less different case", a: "ABC", b: "xyz", expected: -1},
{name: "greater", a: "xyz", b: "abc", expected: 1},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := tt.a.Compare(tt.b); got != tt.expected {
t.Errorf("Compare(%q, %q) = %v, want %v", tt.a, tt.b, got, tt.expected)
}
})
}
}
func TestCIString_Less(t *testing.T) {
if !CIString("abc").Less("xyz") {
t.Error("expected abc < xyz")
}
if CIString("xyz").Less("abc") {
t.Error("expected xyz not < abc")
}
if CIString("ABC").Less("abc") {
t.Error("expected ABC not < abc (equal ignoring case)")
}
}
func TestCIString_Sort(t *testing.T) {
vals := []CIString{"banana", "Apple", "cherry", "apple"}
sort.Slice(vals, func(i, j int) bool { return vals[i].Less(vals[j]) })
// After a case-insensitive sort, "Apple"/"apple" must be adjacent and first,
// followed by banana then cherry.
if !vals[0].EqualString("apple") || !vals[1].EqualString("apple") {
t.Errorf("expected the two apple variants first, got %v", vals)
}
if !vals[2].EqualString("banana") {
t.Errorf("expected banana third, got %v", vals)
}
if !vals[3].EqualString("cherry") {
t.Errorf("expected cherry fourth, got %v", vals)
}
}
func TestLCString_Scan(t *testing.T) {
tests := []struct {
name string
input interface{}
expected LCString
}{
{name: "mixed case string", input: "MixedCase", expected: "mixedcase"},
{name: "bytes mixed case", input: []byte("FromBytes"), expected: "frombytes"},
{name: "nil value", input: nil, expected: ""},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
var s LCString
if err := s.Scan(tt.input); err != nil {
t.Fatalf("Scan failed: %v", err)
}
if s != tt.expected {
t.Errorf("expected %q, got %q", tt.expected, s)
}
})
}
}
func TestLCString_Scan_InvalidType(t *testing.T) {
var s LCString
if err := s.Scan(123); err == nil {
t.Fatal("expected error scanning int into LCString, got nil")
}
}
func TestLCString_Value(t *testing.T) {
s := LCString("MixedCase")
v, err := s.Value()
if err != nil {
t.Fatalf("Value failed: %v", err)
}
if v != "mixedcase" {
t.Errorf("expected %q, got %q", "mixedcase", v)
}
}
func TestLCString_String(t *testing.T) {
s := LCString("MixedCase")
if s.String() != "mixedcase" {
t.Errorf("expected %q, got %q", "mixedcase", s.String())
}
}
func TestLCString_Equal(t *testing.T) {
if !LCString("ABC").Equal(LCString("abc")) {
t.Error("expected ABC and abc to be equal")
}
if LCString("ABC").Equal(LCString("xyz")) {
t.Error("expected ABC and xyz to not be equal")
}
}
func TestLCString_EqualString(t *testing.T) {
if !LCString("ABC").EqualString("abc") {
t.Error("expected EqualString to match case-insensitively")
}
}
func TestUCString_Scan(t *testing.T) {
tests := []struct {
name string
input interface{}
expected UCString
}{
{name: "mixed case string", input: "MixedCase", expected: "MIXEDCASE"},
{name: "bytes mixed case", input: []byte("FromBytes"), expected: "FROMBYTES"},
{name: "nil value", input: nil, expected: ""},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
var s UCString
if err := s.Scan(tt.input); err != nil {
t.Fatalf("Scan failed: %v", err)
}
if s != tt.expected {
t.Errorf("expected %q, got %q", tt.expected, s)
}
})
}
}
func TestUCString_Scan_InvalidType(t *testing.T) {
var s UCString
if err := s.Scan(123); err == nil {
t.Fatal("expected error scanning int into UCString, got nil")
}
}
func TestUCString_Value(t *testing.T) {
s := UCString("MixedCase")
v, err := s.Value()
if err != nil {
t.Fatalf("Value failed: %v", err)
}
if v != "MIXEDCASE" {
t.Errorf("expected %q, got %q", "MIXEDCASE", v)
}
}
func TestUCString_String(t *testing.T) {
s := UCString("MixedCase")
if s.String() != "MIXEDCASE" {
t.Errorf("expected %q, got %q", "MIXEDCASE", s.String())
}
}
func TestUCString_Equal(t *testing.T) {
if !UCString("ABC").Equal(UCString("abc")) {
t.Error("expected ABC and abc to be equal")
}
if UCString("ABC").Equal(UCString("xyz")) {
t.Error("expected ABC and xyz to not be equal")
}
}
func TestUCString_EqualString(t *testing.T) {
if !UCString("ABC").EqualString("abc") {
t.Error("expected EqualString to match case-insensitively")
}
}
// TestLCString_MarshalJSON_NotFromDB verifies a value constructed directly
// in Go (never passed through Scan) still normalizes on JSON marshal.
func TestLCString_MarshalJSON_NotFromDB(t *testing.T) {
s := LCString("MixedCase")
b, err := json.Marshal(s)
if err != nil {
t.Fatalf("Marshal failed: %v", err)
}
if string(b) != `"mixedcase"` {
t.Errorf("expected %s, got %s", `"mixedcase"`, b)
}
}
func TestLCString_UnmarshalJSON(t *testing.T) {
var s LCString
if err := json.Unmarshal([]byte(`"MixedCase"`), &s); err != nil {
t.Fatalf("Unmarshal failed: %v", err)
}
if s != "mixedcase" {
t.Errorf("expected %q, got %q", "mixedcase", s)
}
}
func TestLCString_JSON_StructField(t *testing.T) {
type wrapper struct {
Code LCString `json:"code"`
}
in := wrapper{Code: "MixedCase"}
b, err := json.Marshal(in)
if err != nil {
t.Fatalf("Marshal failed: %v", err)
}
if string(b) != `{"code":"mixedcase"}` {
t.Errorf("expected %s, got %s", `{"code":"mixedcase"}`, b)
}
var out wrapper
if err := json.Unmarshal([]byte(`{"code":"AnotherMixedCase"}`), &out); err != nil {
t.Fatalf("Unmarshal failed: %v", err)
}
if out.Code != "anothermixedcase" {
t.Errorf("expected %q, got %q", "anothermixedcase", out.Code)
}
}
// TestUCString_MarshalJSON_NotFromDB verifies a value constructed directly
// in Go (never passed through Scan) still normalizes on JSON marshal.
func TestUCString_MarshalJSON_NotFromDB(t *testing.T) {
s := UCString("MixedCase")
b, err := json.Marshal(s)
if err != nil {
t.Fatalf("Marshal failed: %v", err)
}
if string(b) != `"MIXEDCASE"` {
t.Errorf("expected %s, got %s", `"MIXEDCASE"`, b)
}
}
func TestUCString_UnmarshalJSON(t *testing.T) {
var s UCString
if err := json.Unmarshal([]byte(`"MixedCase"`), &s); err != nil {
t.Fatalf("Unmarshal failed: %v", err)
}
if s != "MIXEDCASE" {
t.Errorf("expected %q, got %q", "MIXEDCASE", s)
}
}
func TestUCString_JSON_StructField(t *testing.T) {
type wrapper struct {
Code UCString `json:"code"`
}
in := wrapper{Code: "MixedCase"}
b, err := json.Marshal(in)
if err != nil {
t.Fatalf("Marshal failed: %v", err)
}
if string(b) != `{"code":"MIXEDCASE"}` {
t.Errorf("expected %s, got %s", `{"code":"MIXEDCASE"}`, b)
}
var out wrapper
if err := json.Unmarshal([]byte(`{"code":"AnotherMixedCase"}`), &out); err != nil {
t.Fatalf("Unmarshal failed: %v", err)
}
if out.Code != "ANOTHERMIXEDCASE" {
t.Errorf("expected %q, got %q", "ANOTHERMIXEDCASE", out.Code)
}
}
// TestCIString_JSON_PreservesCase confirms CIString needs no custom JSON
// methods: it should never normalize case, only its DB Value/Scan and the
// Equal/EqualString comparisons apply case-insensitivity.
func TestCIString_JSON_PreservesCase(t *testing.T) {
s := CIString("MixedCase")
b, err := json.Marshal(s)
if err != nil {
t.Fatalf("Marshal failed: %v", err)
}
if string(b) != `"MixedCase"` {
t.Errorf("expected %s, got %s", `"MixedCase"`, b)
}
var out CIString
if err := json.Unmarshal([]byte(`"AnotherMixedCase"`), &out); err != nil {
t.Fatalf("Unmarshal failed: %v", err)
}
if out != "AnotherMixedCase" {
t.Errorf("expected case to be preserved, got %q", out)
}
}
+24
View File
@@ -97,3 +97,27 @@ func IsJSONType(t reflect.Type) bool {
n, ok := SQLTypeName(t)
return ok && (n == "jsonb" || n == "json")
}
// UnwrapKind returns the reflect.Kind to use when reasoning about a column's
// comparability (numeric vs. string vs. other) for filter building. Plain Go
// types return their own Kind unchanged. spectypes.SqlNull[T] wrappers (and
// types that embed one, such as SqlTimeStamp/SqlDate/SqlTime) always report
// reflect.Struct for their own Kind even when T is an int64 or string, which
// would otherwise make numeric/text columns look "complex" and force an
// unnecessary CAST(... AS TEXT) that defeats native column indexes. For those
// wrappers, UnwrapKind returns the Kind of the wrapped value T instead.
func UnwrapKind(t reflect.Type) reflect.Kind {
for t != nil && t.Kind() == reflect.Pointer {
t = t.Elem()
}
if t == nil {
return reflect.Invalid
}
if t.Kind() != reflect.Struct || t.PkgPath() != pkgPath {
return t.Kind()
}
if f, ok := t.FieldByName("Val"); ok {
return f.Type.Kind()
}
return t.Kind()
}
+5
View File
@@ -863,6 +863,11 @@ func (h *Handler) buildFilterCondition(filter common.FilterOption, model interfa
op := strings.ToLower(filter.Operator)
if op == "like" || op == "ilike" {
operatorSQL := h.getOperatorSQL(filter.Operator)
// citext columns are already case-insensitive; casting to TEXT would
// switch to case-sensitive matching and defeat a citext index.
if reflection.IsCitextColumn(model, filter.Column) {
return fmt.Sprintf("%s %s ?", filter.Column, operatorSQL), []interface{}{filter.Value}
}
return fmt.Sprintf("CAST(%s AS TEXT) %s ?", filter.Column, operatorSQL), []interface{}{filter.Value}
}
operatorSQL := h.getOperatorSQL(filter.Operator)
+7
View File
@@ -1,5 +1,12 @@
# @warkypublic/resolvespec-js
## 1.0.2
### Patch Changes
- b587cbd: Forward custom ClientConfig headers on every ResolveSpec and HeaderSpec request. Merge headers case-insensitively and isolate cached clients by URL and effective headers, including authentication and tenant headers.
- 7f8982f: fix: added headers and few fixes
## 1.0.1
### Patch Changes
+23 -1
View File
@@ -28,7 +28,7 @@ import { ResolveSpecClient, getResolveSpecClient } from '@warkypublic/resolvespe
// Class instantiation
const client = new ResolveSpecClient({ baseUrl: 'http://localhost:3000', token: 'your-token' });
// Or singleton factory (returns cached instance per baseUrl)
// Or singleton factory (returns cached instance per baseUrl and effective headers)
const client = getResolveSpecClient({ baseUrl: 'http://localhost:3000', token: 'your-token' });
// Read with filters, sort, pagination
@@ -211,3 +211,25 @@ pnpm run lint # eslint
## License
MIT
### Custom HTTP headers
Both `ResolveSpecClient` and `HeaderSpecClient` (including their factory functions)
accept `headers` in `ClientConfig` and send them on every HTTP request:
```typescript
const client = new ResolveSpecClient({
baseUrl: 'http://localhost:3000',
token: 'your-token',
headers: { 'X-Tenant': 'acme' },
});
```
Header names are merged case-insensitively. Custom headers override the default
`Content-Type`; a supplied `token` overrides custom `Authorization`, and HeaderSpec
query options override matching custom query headers. Without a token, custom
`Authorization` is preserved. Configuration is copied at construction; create or
retrieve a client with new configuration to change headers. Factory clients are
cached by URL and effective headers, keeping different tenants and tokens separate.
Grid adapters must forward `dataSourceOptions.headers` to this `headers` option.
+1 -1
View File
File diff suppressed because one or more lines are too long
+5 -366
View File
@@ -1,366 +1,5 @@
export declare interface APIError {
code: string;
message: string;
details?: any;
detail?: string;
}
export declare interface APIResponse<T = any> {
success: boolean;
data: T;
metadata?: Metadata;
error?: APIError;
}
/**
* Build HTTP headers from Options, matching Go's restheadspec handler conventions.
*
* Header mapping:
* - X-Select-Fields: comma-separated columns
* - X-Not-Select-Fields: comma-separated omit_columns
* - X-FieldFilter-{col}: exact match (eq)
* - X-SearchOp-{operator}-{col}: AND filter
* - X-SearchOr-{operator}-{col}: OR filter
* - X-Sort: +col (asc), -col (desc)
* - X-Limit, X-Offset: pagination
* - X-Cursor-Forward, X-Cursor-Backward: cursor pagination
* - X-Preload: RelationName:field1,field2 pipe-separated
* - X-Fetch-RowNumber: row number fetch
* - X-CQL-SEL-{col}: computed columns
* - X-Custom-SQL-W: custom operators (AND)
*/
export declare function buildHeaders(options: Options): Record<string, string>;
export declare interface ClientConfig {
baseUrl: string;
token?: string;
}
export declare interface Column {
name: string;
type: string;
is_nullable: boolean;
is_primary: boolean;
is_unique: boolean;
has_index: boolean;
}
export declare interface ComputedColumn {
name: string;
expression: string;
}
export declare type ConnectionState = 'connecting' | 'connected' | 'disconnecting' | 'disconnected' | 'reconnecting';
export declare interface CustomOperator {
name: string;
sql: string;
}
/**
* Decode a header value that may be base64 encoded with ZIP_ or __ prefix.
*/
export declare function decodeHeaderValue(value: string): string;
/**
* Encode a value with base64 and ZIP_ prefix for complex header values.
*/
export declare function encodeHeaderValue(value: string): string;
export declare interface FilterOption {
column: string;
operator: Operator | string;
value: any;
logic_operator?: 'AND' | 'OR';
}
export declare function getHeaderSpecClient(config: ClientConfig): HeaderSpecClient;
export declare function getResolveSpecClient(config: ClientConfig): ResolveSpecClient;
export declare function getWebSocketClient(config: WebSocketClientConfig): WebSocketClient;
/**
* HeaderSpec REST client.
* Sends query options via HTTP headers instead of request body, matching the Go restheadspec handler.
*
* HTTP methods: GET=read, POST=create, PUT=update, DELETE=delete
*/
export declare class HeaderSpecClient {
private config;
constructor(config: ClientConfig);
private buildUrl;
private baseHeaders;
private fetchWithError;
read<T = any>(schema: string, entity: string, id?: string, options?: Options): Promise<APIResponse<T>>;
create<T = any>(schema: string, entity: string, data: any, options?: Options): Promise<APIResponse<T>>;
update<T = any>(schema: string, entity: string, id: string, data: any, options?: Options): Promise<APIResponse<T>>;
delete(schema: string, entity: string, id: string): Promise<APIResponse<void>>;
}
export declare type MessageType = 'request' | 'response' | 'notification' | 'subscription' | 'error' | 'ping' | 'pong';
export declare interface Metadata {
total: number;
count: number;
filtered: number;
limit: number;
offset: number;
row_number?: number;
}
export declare type Operation = 'read' | 'create' | 'update' | 'delete';
export declare type Operator = 'eq' | 'neq' | 'gt' | 'gte' | 'lt' | 'lte' | 'like' | 'ilike' | 'in' | 'contains' | 'startswith' | 'endswith' | 'between' | 'between_inclusive' | 'is_null' | 'is_not_null';
export declare interface Options {
preload?: PreloadOption[];
columns?: string[];
omit_columns?: string[];
filters?: FilterOption[];
sort?: SortOption[];
limit?: number;
offset?: number;
customOperators?: CustomOperator[];
computedColumns?: ComputedColumn[];
parameters?: Parameter[];
cursor_forward?: string;
cursor_backward?: string;
fetch_row_number?: string;
}
export declare interface Parameter {
name: string;
value: string;
sequence?: number;
}
export declare interface PreloadOption {
relation: string;
table_name?: string;
columns?: string[];
omit_columns?: string[];
sort?: SortOption[];
filters?: FilterOption[];
where?: string;
limit?: number;
offset?: number;
updatable?: boolean;
computed_ql?: Record<string, string>;
recursive?: boolean;
primary_key?: string;
related_key?: string;
foreign_key?: string;
recursive_child_key?: string;
sql_joins?: string[];
join_aliases?: string[];
}
export declare interface RequestBody {
operation: Operation;
id?: number | string | string[];
data?: any | any[];
options?: Options;
}
export declare class ResolveSpecClient {
private config;
constructor(config: ClientConfig);
private buildUrl;
private baseHeaders;
private fetchWithError;
getMetadata(schema: string, entity: string): Promise<APIResponse<TableMetadata>>;
read<T = any>(schema: string, entity: string, id?: number | string | string[], options?: Options): Promise<APIResponse<T>>;
create<T = any>(schema: string, entity: string, data: any | any[], options?: Options): Promise<APIResponse<T>>;
update<T = any>(schema: string, entity: string, data: any | any[], id?: number | string | string[], options?: Options): Promise<APIResponse<T>>;
delete(schema: string, entity: string, id: number | string): Promise<APIResponse<void>>;
}
export declare type SortDirection = 'asc' | 'desc' | 'ASC' | 'DESC';
export declare interface SortOption {
column: string;
direction: SortDirection;
}
export declare interface Subscription {
id: string;
entity: string;
schema?: string;
options?: WSOptions;
callback?: (notification: WSNotificationMessage) => void;
}
export declare interface SubscriptionOptions {
filters?: FilterOption[];
onNotification?: (notification: WSNotificationMessage) => void;
}
export declare interface TableMetadata {
schema: string;
table: string;
columns: Column[];
relations: string[];
}
export declare class WebSocketClient {
private ws;
private config;
private messageHandlers;
private subscriptions;
private eventListeners;
private state;
private reconnectAttempts;
private reconnectTimer;
private heartbeatTimer;
private isManualClose;
constructor(config: WebSocketClientConfig);
connect(): Promise<void>;
disconnect(): void;
request<T = any>(operation: WSOperation, entity: string, options?: {
schema?: string;
record_id?: string;
data?: any;
options?: WSOptions;
}): Promise<T>;
read<T = any>(entity: string, options?: {
schema?: string;
record_id?: string;
filters?: FilterOption[];
columns?: string[];
sort?: SortOption[];
preload?: PreloadOption[];
limit?: number;
offset?: number;
}): Promise<T>;
create<T = any>(entity: string, data: any, options?: {
schema?: string;
}): Promise<T>;
update<T = any>(entity: string, id: string, data: any, options?: {
schema?: string;
}): Promise<T>;
delete(entity: string, id: string, options?: {
schema?: string;
}): Promise<void>;
meta<T = any>(entity: string, options?: {
schema?: string;
}): Promise<T>;
subscribe(entity: string, callback: (notification: WSNotificationMessage) => void, options?: {
schema?: string;
filters?: FilterOption[];
}): Promise<string>;
unsubscribe(subscriptionId: string): Promise<void>;
getSubscriptions(): Subscription[];
getState(): ConnectionState;
isConnected(): boolean;
on<K extends keyof WebSocketClientEvents>(event: K, callback: WebSocketClientEvents[K]): void;
off<K extends keyof WebSocketClientEvents>(event: K): void;
private handleMessage;
private handleResponse;
private handleNotification;
private send;
private startHeartbeat;
private stopHeartbeat;
private setState;
private ensureConnected;
private emit;
private log;
}
export declare interface WebSocketClientConfig {
url: string;
reconnect?: boolean;
reconnectInterval?: number;
maxReconnectAttempts?: number;
heartbeatInterval?: number;
debug?: boolean;
}
export declare interface WebSocketClientEvents {
connect: () => void;
disconnect: (event: CloseEvent) => void;
error: (error: Error) => void;
message: (message: WSMessage) => void;
stateChange: (state: ConnectionState) => void;
}
export declare interface WSErrorInfo {
code: string;
message: string;
details?: Record<string, any>;
}
export declare interface WSMessage {
id?: string;
type: MessageType;
operation?: WSOperation;
schema?: string;
entity?: string;
record_id?: string;
data?: any;
options?: WSOptions;
subscription_id?: string;
success?: boolean;
error?: WSErrorInfo;
metadata?: Record<string, any>;
timestamp?: string;
}
export declare interface WSNotificationMessage {
type: 'notification';
operation: WSOperation;
subscription_id: string;
schema?: string;
entity: string;
data: any;
timestamp: string;
}
export declare type WSOperation = 'read' | 'create' | 'update' | 'delete' | 'subscribe' | 'unsubscribe' | 'meta';
export declare interface WSOptions {
filters?: FilterOption[];
columns?: string[];
omit_columns?: string[];
preload?: PreloadOption[];
sort?: SortOption[];
limit?: number;
offset?: number;
parameters?: Parameter[];
cursor_forward?: string;
cursor_backward?: string;
fetch_row_number?: string;
}
export declare interface WSRequestMessage {
id: string;
type: 'request';
operation: WSOperation;
schema?: string;
entity: string;
record_id?: string;
data?: any;
options?: WSOptions;
}
export declare interface WSResponseMessage {
id: string;
type: 'response';
success: boolean;
data?: any;
error?: WSErrorInfo;
metadata?: Record<string, any>;
timestamp: string;
}
export declare interface WSSubscriptionMessage {
id: string;
type: 'subscription';
operation: 'subscribe' | 'unsubscribe';
schema?: string;
entity: string;
options?: WSOptions;
subscription_id?: string;
}
export { }
export * from './common';
export * from './resolvespec';
export * from './websocketspec';
export * from './headerspec';
//# sourceMappingURL=index.d.ts.map
+426 -463
View File
@@ -1,469 +1,432 @@
import { v4 as l } from "uuid";
const d = /* @__PURE__ */ new Map();
function E(n) {
const e = n.baseUrl;
let t = d.get(e);
return t || (t = new g(n), d.set(e, t)), t;
import { v4 as e } from "uuid";
import { b64DecodeUnicode as t, b64EncodeUnicode as n } from "@warkypublic/artemis-kit/base64";
//#region src/common/http.ts
function r(...e) {
let t = {};
for (let n of e) for (let [e, r] of Object.entries(n)) {
for (let n of Object.keys(t)) n.toLowerCase() === e.toLowerCase() && delete t[n];
Object.defineProperty(t, e, {
value: r,
enumerable: !0,
configurable: !0,
writable: !0
});
}
return t;
}
class g {
constructor(e) {
this.config = e;
}
buildUrl(e, t, s) {
let r = `${this.config.baseUrl}/${e}/${t}`;
return s && (r += `/${s}`), r;
}
baseHeaders() {
const e = {
"Content-Type": "application/json"
};
return this.config.token && (e.Authorization = `Bearer ${this.config.token}`), e;
}
async fetchWithError(e, t) {
const s = await fetch(e, t), r = await s.json();
if (!s.ok)
throw new Error(r.error?.message || "An error occurred");
return r;
}
async getMetadata(e, t) {
const s = this.buildUrl(e, t);
return this.fetchWithError(s, {
method: "GET",
headers: this.baseHeaders()
});
}
async read(e, t, s, r) {
const i = typeof s == "number" || typeof s == "string" ? String(s) : void 0, a = this.buildUrl(e, t, i), c = {
operation: "read",
id: Array.isArray(s) ? s : void 0,
options: r
};
return this.fetchWithError(a, {
method: "POST",
headers: this.baseHeaders(),
body: JSON.stringify(c)
});
}
async create(e, t, s, r) {
const i = this.buildUrl(e, t), a = {
operation: "create",
data: s,
options: r
};
return this.fetchWithError(i, {
method: "POST",
headers: this.baseHeaders(),
body: JSON.stringify(a)
});
}
async update(e, t, s, r, i) {
const a = typeof r == "number" || typeof r == "string" ? String(r) : void 0, c = this.buildUrl(e, t, a), o = {
operation: "update",
id: Array.isArray(r) ? r : void 0,
data: s,
options: i
};
return this.fetchWithError(c, {
method: "POST",
headers: this.baseHeaders(),
body: JSON.stringify(o)
});
}
async delete(e, t, s) {
const r = this.buildUrl(e, t, String(s)), i = {
operation: "delete"
};
return this.fetchWithError(r, {
method: "POST",
headers: this.baseHeaders(),
body: JSON.stringify(i)
});
}
function i(e) {
return r({ "Content-Type": "application/json" }, e.headers ?? {}, e.token ? { Authorization: `Bearer ${e.token}` } : {});
}
const f = /* @__PURE__ */ new Map();
function _(n) {
const e = n.url;
let t = f.get(e);
return t || (t = new p(n), f.set(e, t)), t;
function a(e) {
let t = Object.entries(i(e)).map(([e, t]) => [e.toLowerCase(), t]).sort(([e], [t]) => e.localeCompare(t));
return JSON.stringify([e.baseUrl, t]);
}
class p {
constructor(e) {
this.ws = null, this.messageHandlers = /* @__PURE__ */ new Map(), this.subscriptions = /* @__PURE__ */ new Map(), this.eventListeners = {}, this.state = "disconnected", this.reconnectAttempts = 0, this.reconnectTimer = null, this.heartbeatTimer = null, this.isManualClose = !1, this.config = {
url: e.url,
reconnect: e.reconnect ?? !0,
reconnectInterval: e.reconnectInterval ?? 3e3,
maxReconnectAttempts: e.maxReconnectAttempts ?? 10,
heartbeatInterval: e.heartbeatInterval ?? 3e4,
debug: e.debug ?? !1
};
}
async connect() {
if (this.ws?.readyState === WebSocket.OPEN) {
this.log("Already connected");
return;
}
return this.isManualClose = !1, this.setState("connecting"), new Promise((e, t) => {
try {
this.ws = new WebSocket(this.config.url), this.ws.onopen = () => {
this.log("Connected to WebSocket server"), this.setState("connected"), this.reconnectAttempts = 0, this.startHeartbeat(), this.emit("connect"), e();
}, this.ws.onmessage = (s) => {
this.handleMessage(s.data);
}, this.ws.onerror = (s) => {
this.log("WebSocket error:", s);
const r = new Error("WebSocket connection error");
this.emit("error", r), t(r);
}, this.ws.onclose = (s) => {
this.log("WebSocket closed:", s.code, s.reason), this.stopHeartbeat(), this.setState("disconnected"), this.emit("disconnect", s), this.config.reconnect && !this.isManualClose && this.reconnectAttempts < this.config.maxReconnectAttempts && (this.reconnectAttempts++, this.log(`Reconnection attempt ${this.reconnectAttempts}/${this.config.maxReconnectAttempts}`), this.setState("reconnecting"), this.reconnectTimer = setTimeout(() => {
this.connect().catch((r) => {
this.log("Reconnection failed:", r);
});
}, this.config.reconnectInterval));
};
} catch (s) {
t(s);
}
});
}
disconnect() {
this.isManualClose = !0, this.reconnectTimer && (clearTimeout(this.reconnectTimer), this.reconnectTimer = null), this.stopHeartbeat(), this.ws && (this.setState("disconnecting"), this.ws.close(), this.ws = null), this.setState("disconnected"), this.messageHandlers.clear();
}
async request(e, t, s) {
this.ensureConnected();
const r = l(), i = {
id: r,
type: "request",
operation: e,
entity: t,
schema: s?.schema,
record_id: s?.record_id,
data: s?.data,
options: s?.options
};
return new Promise((a, c) => {
this.messageHandlers.set(r, (o) => {
o.success ? a(o.data) : c(new Error(o.error?.message || "Request failed"));
}), this.send(i), setTimeout(() => {
this.messageHandlers.has(r) && (this.messageHandlers.delete(r), c(new Error("Request timeout")));
}, 3e4);
});
}
async read(e, t) {
return this.request("read", e, {
schema: t?.schema,
record_id: t?.record_id,
options: {
filters: t?.filters,
columns: t?.columns,
sort: t?.sort,
preload: t?.preload,
limit: t?.limit,
offset: t?.offset
}
});
}
async create(e, t, s) {
return this.request("create", e, {
schema: s?.schema,
data: t
});
}
async update(e, t, s, r) {
return this.request("update", e, {
schema: r?.schema,
record_id: t,
data: s
});
}
async delete(e, t, s) {
await this.request("delete", e, {
schema: s?.schema,
record_id: t
});
}
async meta(e, t) {
return this.request("meta", e, {
schema: t?.schema
});
}
async subscribe(e, t, s) {
this.ensureConnected();
const r = l(), i = {
id: r,
type: "subscription",
operation: "subscribe",
entity: e,
schema: s?.schema,
options: {
filters: s?.filters
}
};
return new Promise((a, c) => {
this.messageHandlers.set(r, (o) => {
if (o.success && o.data?.subscription_id) {
const h = o.data.subscription_id;
this.subscriptions.set(h, {
id: h,
entity: e,
schema: s?.schema,
options: { filters: s?.filters },
callback: t
}), this.log(`Subscribed to ${e} with ID: ${h}`), a(h);
} else
c(new Error(o.error?.message || "Subscription failed"));
}), this.send(i), setTimeout(() => {
this.messageHandlers.has(r) && (this.messageHandlers.delete(r), c(new Error("Subscription timeout")));
}, 1e4);
});
}
async unsubscribe(e) {
this.ensureConnected();
const t = l(), s = {
id: t,
type: "subscription",
operation: "unsubscribe",
subscription_id: e
};
return new Promise((r, i) => {
this.messageHandlers.set(t, (a) => {
a.success ? (this.subscriptions.delete(e), this.log(`Unsubscribed from ${e}`), r()) : i(new Error(a.error?.message || "Unsubscribe failed"));
}), this.send(s), setTimeout(() => {
this.messageHandlers.has(t) && (this.messageHandlers.delete(t), i(new Error("Unsubscribe timeout")));
}, 1e4);
});
}
getSubscriptions() {
return Array.from(this.subscriptions.values());
}
getState() {
return this.state;
}
isConnected() {
return this.ws?.readyState === WebSocket.OPEN;
}
on(e, t) {
this.eventListeners[e] = t;
}
off(e) {
delete this.eventListeners[e];
}
// Private methods
handleMessage(e) {
try {
const t = JSON.parse(e);
switch (this.log("Received message:", t), this.emit("message", t), t.type) {
case "response":
this.handleResponse(t);
break;
case "notification":
this.handleNotification(t);
break;
case "pong":
break;
default:
this.log("Unknown message type:", t.type);
}
} catch (t) {
this.log("Error parsing message:", t);
}
}
handleResponse(e) {
const t = this.messageHandlers.get(e.id);
t && (t(e), this.messageHandlers.delete(e.id));
}
handleNotification(e) {
const t = this.subscriptions.get(e.subscription_id);
t?.callback && t.callback(e);
}
send(e) {
if (!this.ws || this.ws.readyState !== WebSocket.OPEN)
throw new Error("WebSocket is not connected");
const t = JSON.stringify(e);
this.log("Sending message:", e), this.ws.send(t);
}
startHeartbeat() {
this.heartbeatTimer || (this.heartbeatTimer = setInterval(() => {
if (this.isConnected()) {
const e = {
id: l(),
type: "ping"
};
this.send(e);
}
}, this.config.heartbeatInterval));
}
stopHeartbeat() {
this.heartbeatTimer && (clearInterval(this.heartbeatTimer), this.heartbeatTimer = null);
}
setState(e) {
this.state !== e && (this.state = e, this.emit("stateChange", e));
}
ensureConnected() {
if (!this.isConnected())
throw new Error("WebSocket is not connected. Call connect() first.");
}
emit(e, ...t) {
const s = this.eventListeners[e];
s && s(...t);
}
log(...e) {
this.config.debug && console.log("[WebSocketClient]", ...e);
}
//#endregion
//#region src/resolvespec/client.ts
var o = /* @__PURE__ */ new Map();
function s(e) {
let t = a(e), n = o.get(t);
return n || (n = new c(e), o.set(t, n)), n;
}
function v(n) {
return typeof btoa == "function" ? "ZIP_" + btoa(n) : "ZIP_" + Buffer.from(n, "utf-8").toString("base64");
var c = class {
constructor(e) {
this.config = {
...e,
headers: { ...e.headers }
};
}
buildUrl(e, t, n) {
let r = `${this.config.baseUrl}/${e}/${t}`;
return n && (r += `/${n}`), r;
}
baseHeaders() {
return i(this.config);
}
async fetchWithError(e, t) {
let n = await fetch(e, t), r = await n.json();
if (!n.ok) throw Error(r.error?.message || "An error occurred");
return r;
}
async getMetadata(e, t) {
let n = this.buildUrl(e, t);
return this.fetchWithError(n, {
method: "GET",
headers: this.baseHeaders()
});
}
async read(e, t, n, r) {
let i = typeof n == "number" || typeof n == "string" ? String(n) : void 0, a = this.buildUrl(e, t, i), o = {
operation: "read",
id: Array.isArray(n) ? n : void 0,
options: r
};
return this.fetchWithError(a, {
method: "POST",
headers: this.baseHeaders(),
body: JSON.stringify(o)
});
}
async create(e, t, n, r) {
let i = this.buildUrl(e, t), a = {
operation: "create",
data: n,
options: r
};
return this.fetchWithError(i, {
method: "POST",
headers: this.baseHeaders(),
body: JSON.stringify(a)
});
}
async update(e, t, n, r, i) {
let a = typeof r == "number" || typeof r == "string" ? String(r) : void 0, o = this.buildUrl(e, t, a), s = {
operation: "update",
id: Array.isArray(r) ? r : void 0,
data: n,
options: i
};
return this.fetchWithError(o, {
method: "POST",
headers: this.baseHeaders(),
body: JSON.stringify(s)
});
}
async delete(e, t, n) {
let r = this.buildUrl(e, t, String(n));
return this.fetchWithError(r, {
method: "POST",
headers: this.baseHeaders(),
body: JSON.stringify({ operation: "delete" })
});
}
}, l = /* @__PURE__ */ new Map();
function u(e) {
let t = e.url, n = l.get(t);
return n || (n = new d(e), l.set(t, n)), n;
}
function w(n) {
let e = n;
return e.startsWith("ZIP_") ? (e = e.slice(4).replace(/[\n\r ]/g, ""), e = m(e)) : e.startsWith("__") && (e = e.slice(2).replace(/[\n\r ]/g, ""), e = m(e)), (e.startsWith("ZIP_") || e.startsWith("__")) && (e = w(e)), e;
}
function m(n) {
return typeof atob == "function" ? atob(n) : Buffer.from(n, "base64").toString("utf-8");
}
function u(n) {
const e = {};
if (n.columns?.length && (e["X-Select-Fields"] = n.columns.join(",")), n.omit_columns?.length && (e["X-Not-Select-Fields"] = n.omit_columns.join(",")), n.filters?.length)
for (const t of n.filters) {
const s = t.logic_operator ?? "AND", r = y(t.operator), i = S(t);
t.operator === "eq" && s === "AND" ? e[`X-FieldFilter-${t.column}`] = i : s === "OR" ? e[`X-SearchOr-${r}-${t.column}`] = i : e[`X-SearchOp-${r}-${t.column}`] = i;
}
if (n.sort?.length) {
const t = n.sort.map((s) => s.direction.toUpperCase() === "DESC" ? `-${s.column}` : `+${s.column}`);
e["X-Sort"] = t.join(",");
}
if (n.limit !== void 0 && (e["X-Limit"] = String(n.limit)), n.offset !== void 0 && (e["X-Offset"] = String(n.offset)), n.cursor_forward && (e["X-Cursor-Forward"] = n.cursor_forward), n.cursor_backward && (e["X-Cursor-Backward"] = n.cursor_backward), n.preload?.length) {
const t = n.preload.map((s) => s.columns?.length ? `${s.relation}:${s.columns.join(",")}` : s.relation);
e["X-Preload"] = t.join("|");
}
if (n.fetch_row_number && (e["X-Fetch-RowNumber"] = n.fetch_row_number), n.computedColumns?.length)
for (const t of n.computedColumns)
e[`X-CQL-SEL-${t.name}`] = t.expression;
if (n.customOperators?.length) {
const t = n.customOperators.map(
(s) => s.sql
);
e["X-Custom-SQL-W"] = t.join(" AND ");
}
return e;
}
function y(n) {
switch (n) {
case "eq":
return "equals";
case "neq":
return "notequals";
case "gt":
return "greaterthan";
case "gte":
return "greaterthanorequal";
case "lt":
return "lessthan";
case "lte":
return "lessthanorequal";
case "like":
case "ilike":
case "contains":
return "contains";
case "startswith":
return "beginswith";
case "endswith":
return "endswith";
case "in":
return "in";
case "between":
return "between";
case "between_inclusive":
return "betweeninclusive";
case "is_null":
return "empty";
case "is_not_null":
return "notempty";
default:
return n;
}
}
function S(n) {
return n.value === null || n.value === void 0 ? "" : Array.isArray(n.value) ? n.value.join(",") : String(n.value);
}
const b = /* @__PURE__ */ new Map();
function C(n) {
const e = n.baseUrl;
let t = b.get(e);
return t || (t = new H(n), b.set(e, t)), t;
}
class H {
constructor(e) {
this.config = e;
}
buildUrl(e, t, s) {
let r = `${this.config.baseUrl}/${e}/${t}`;
return s && (r += `/${s}`), r;
}
baseHeaders() {
const e = {
"Content-Type": "application/json"
};
return this.config.token && (e.Authorization = `Bearer ${this.config.token}`), e;
}
async fetchWithError(e, t) {
const s = await fetch(e, t), r = await s.json();
if (!s.ok)
throw new Error(
r.error?.message || `${s.statusText} (${s.status})`
);
return {
data: r,
success: !0,
error: r.error ? r.error : void 0,
metadata: {
count: s.headers.get("content-range") ? Number(s.headers.get("content-range")?.split("/")[1]) : 0,
total: s.headers.get("content-range") ? Number(s.headers.get("content-range")?.split("/")[1]) : 0,
filtered: s.headers.get("content-range") ? Number(s.headers.get("content-range")?.split("/")[1]) : 0,
offset: s.headers.get("content-range") ? Number(
s.headers.get("content-range")?.split("/")[0].split("-")[0]
) : 0,
limit: s.headers.get("x-limit") ? Number(s.headers.get("x-limit")) : 0
}
};
}
async read(e, t, s, r) {
const i = this.buildUrl(e, t, s), a = r ? u(r) : {};
return this.fetchWithError(i, {
method: "GET",
headers: { ...this.baseHeaders(), ...a }
});
}
async create(e, t, s, r) {
const i = this.buildUrl(e, t), a = r ? u(r) : {};
return this.fetchWithError(i, {
method: "POST",
headers: { ...this.baseHeaders(), ...a },
body: JSON.stringify(s)
});
}
async update(e, t, s, r, i) {
const a = this.buildUrl(e, t, s), c = i ? u(i) : {};
return this.fetchWithError(a, {
method: "PUT",
headers: { ...this.baseHeaders(), ...c },
body: JSON.stringify(r)
});
}
async delete(e, t, s) {
const r = this.buildUrl(e, t, s);
return this.fetchWithError(r, {
method: "DELETE",
headers: this.baseHeaders()
});
}
}
export {
H as HeaderSpecClient,
g as ResolveSpecClient,
p as WebSocketClient,
u as buildHeaders,
w as decodeHeaderValue,
v as encodeHeaderValue,
C as getHeaderSpecClient,
E as getResolveSpecClient,
_ as getWebSocketClient
var d = class {
constructor(e) {
this.ws = null, this.messageHandlers = /* @__PURE__ */ new Map(), this.subscriptions = /* @__PURE__ */ new Map(), this.eventListeners = {}, this.state = "disconnected", this.reconnectAttempts = 0, this.reconnectTimer = null, this.heartbeatTimer = null, this.isManualClose = !1, this.config = {
url: e.url,
reconnect: e.reconnect ?? !0,
reconnectInterval: e.reconnectInterval ?? 3e3,
maxReconnectAttempts: e.maxReconnectAttempts ?? 10,
heartbeatInterval: e.heartbeatInterval ?? 3e4,
debug: e.debug ?? !1
};
}
async connect() {
if (this.ws?.readyState === WebSocket.OPEN) {
this.log("Already connected");
return;
}
return this.isManualClose = !1, this.setState("connecting"), new Promise((e, t) => {
try {
this.ws = new WebSocket(this.config.url), this.ws.onopen = () => {
this.log("Connected to WebSocket server"), this.setState("connected"), this.reconnectAttempts = 0, this.startHeartbeat(), this.emit("connect"), e();
}, this.ws.onmessage = (e) => {
this.handleMessage(e.data);
}, this.ws.onerror = (e) => {
this.log("WebSocket error:", e);
let n = /* @__PURE__ */ Error("WebSocket connection error");
this.emit("error", n), t(n);
}, this.ws.onclose = (e) => {
this.log("WebSocket closed:", e.code, e.reason), this.stopHeartbeat(), this.setState("disconnected"), this.emit("disconnect", e), this.config.reconnect && !this.isManualClose && this.reconnectAttempts < this.config.maxReconnectAttempts && (this.reconnectAttempts++, this.log(`Reconnection attempt ${this.reconnectAttempts}/${this.config.maxReconnectAttempts}`), this.setState("reconnecting"), this.reconnectTimer = setTimeout(() => {
this.connect().catch((e) => {
this.log("Reconnection failed:", e);
});
}, this.config.reconnectInterval));
};
} catch (e) {
t(e);
}
});
}
disconnect() {
this.isManualClose = !0, this.reconnectTimer &&= (clearTimeout(this.reconnectTimer), null), this.stopHeartbeat(), this.ws &&= (this.setState("disconnecting"), this.ws.close(), null), this.setState("disconnected"), this.messageHandlers.clear();
}
async request(t, n, r) {
this.ensureConnected();
let i = e(), a = {
id: i,
type: "request",
operation: t,
entity: n,
schema: r?.schema,
record_id: r?.record_id,
data: r?.data,
options: r?.options
};
return new Promise((e, t) => {
this.messageHandlers.set(i, (n) => {
n.success ? e(n.data) : t(Error(n.error?.message || "Request failed"));
}), this.send(a), setTimeout(() => {
this.messageHandlers.has(i) && (this.messageHandlers.delete(i), t(/* @__PURE__ */ Error("Request timeout")));
}, 3e4);
});
}
async read(e, t) {
return this.request("read", e, {
schema: t?.schema,
record_id: t?.record_id,
options: {
filters: t?.filters,
columns: t?.columns,
sort: t?.sort,
preload: t?.preload,
limit: t?.limit,
offset: t?.offset
}
});
}
async create(e, t, n) {
return this.request("create", e, {
schema: n?.schema,
data: t
});
}
async update(e, t, n, r) {
return this.request("update", e, {
schema: r?.schema,
record_id: t,
data: n
});
}
async delete(e, t, n) {
await this.request("delete", e, {
schema: n?.schema,
record_id: t
});
}
async meta(e, t) {
return this.request("meta", e, { schema: t?.schema });
}
async subscribe(t, n, r) {
this.ensureConnected();
let i = e(), a = {
id: i,
type: "subscription",
operation: "subscribe",
entity: t,
schema: r?.schema,
options: { filters: r?.filters }
};
return new Promise((e, o) => {
this.messageHandlers.set(i, (i) => {
if (i.success && i.data?.subscription_id) {
let a = i.data.subscription_id;
this.subscriptions.set(a, {
id: a,
entity: t,
schema: r?.schema,
options: { filters: r?.filters },
callback: n
}), this.log(`Subscribed to ${t} with ID: ${a}`), e(a);
} else o(Error(i.error?.message || "Subscription failed"));
}), this.send(a), setTimeout(() => {
this.messageHandlers.has(i) && (this.messageHandlers.delete(i), o(/* @__PURE__ */ Error("Subscription timeout")));
}, 1e4);
});
}
async unsubscribe(t) {
this.ensureConnected();
let n = e(), r = {
id: n,
type: "subscription",
operation: "unsubscribe",
subscription_id: t
};
return new Promise((e, i) => {
this.messageHandlers.set(n, (n) => {
n.success ? (this.subscriptions.delete(t), this.log(`Unsubscribed from ${t}`), e()) : i(Error(n.error?.message || "Unsubscribe failed"));
}), this.send(r), setTimeout(() => {
this.messageHandlers.has(n) && (this.messageHandlers.delete(n), i(/* @__PURE__ */ Error("Unsubscribe timeout")));
}, 1e4);
});
}
getSubscriptions() {
return Array.from(this.subscriptions.values());
}
getState() {
return this.state;
}
isConnected() {
return this.ws?.readyState === WebSocket.OPEN;
}
on(e, t) {
this.eventListeners[e] = t;
}
off(e) {
delete this.eventListeners[e];
}
handleMessage(e) {
try {
let t = JSON.parse(e);
switch (this.log("Received message:", t), this.emit("message", t), t.type) {
case "response":
this.handleResponse(t);
break;
case "notification":
this.handleNotification(t);
break;
case "pong": break;
default: this.log("Unknown message type:", t.type);
}
} catch (e) {
this.log("Error parsing message:", e);
}
}
handleResponse(e) {
let t = this.messageHandlers.get(e.id);
t && (t(e), this.messageHandlers.delete(e.id));
}
handleNotification(e) {
let t = this.subscriptions.get(e.subscription_id);
t?.callback && t.callback(e);
}
send(e) {
if (!this.ws || this.ws.readyState !== WebSocket.OPEN) throw Error("WebSocket is not connected");
let t = JSON.stringify(e);
this.log("Sending message:", e), this.ws.send(t);
}
startHeartbeat() {
this.heartbeatTimer ||= setInterval(() => {
if (this.isConnected()) {
let t = {
id: e(),
type: "ping"
};
this.send(t);
}
}, this.config.heartbeatInterval);
}
stopHeartbeat() {
this.heartbeatTimer &&= (clearInterval(this.heartbeatTimer), null);
}
setState(e) {
this.state !== e && (this.state = e, this.emit("stateChange", e));
}
ensureConnected() {
if (!this.isConnected()) throw Error("WebSocket is not connected. Call connect() first.");
}
emit(e, ...t) {
let n = this.eventListeners[e];
n && n(...t);
}
log(...e) {
this.config.debug && console.log("[WebSocketClient]", ...e);
}
};
//#endregion
//#region src/headerspec/client.ts
function f(e) {
return "ZIP_" + n(e);
}
function p(e) {
let t = e;
return t.startsWith("ZIP_") ? (t = t.slice(4).replace(/[\n\r ]/g, ""), t = m(t)) : t.startsWith("__") && (t = t.slice(2).replace(/[\n\r ]/g, ""), t = m(t)), (t.startsWith("ZIP_") || t.startsWith("__")) && (t = p(t)), t;
}
function m(e) {
return t(e);
}
function h(e) {
let t = {};
if (e.columns?.length && (t["X-Select-Fields"] = e.columns.join(",")), e.omit_columns?.length && (t["X-Not-Select-Fields"] = e.omit_columns.join(",")), e.filters?.length) for (let n of e.filters) {
let e = n.logic_operator ?? "AND", r = g(n.operator), i = _(n);
n.operator === "eq" && e === "AND" ? t[`X-FieldFilter-${n.column}`] = i : e === "OR" ? t[`X-SearchOr-${r}-${n.column}`] = i : t[`X-SearchOp-${r}-${n.column}`] = i;
}
if (e.sort?.length && (t["X-Sort"] = e.sort.map((e) => e.direction.toUpperCase() === "DESC" ? `-${e.column}` : `+${e.column}`).join(",")), e.limit !== void 0 && (t["X-Limit"] = String(e.limit)), e.offset !== void 0 && (t["X-Offset"] = String(e.offset)), e.cursor_forward && (t["X-Cursor-Forward"] = e.cursor_forward), e.cursor_backward && (t["X-Cursor-Backward"] = e.cursor_backward), e.preload?.length && (t["X-Preload"] = e.preload.map((e) => e.columns?.length ? `${e.relation}:${e.columns.join(",")}` : e.relation).join("|")), e.fetch_row_number && (t["X-Fetch-RowNumber"] = e.fetch_row_number), e.computedColumns?.length) for (let n of e.computedColumns) t[`X-CQL-SEL-${n.name}`] = n.expression;
return e.customOperators?.length && (t["X-Custom-SQL-W"] = e.customOperators.map((e) => e.sql).join(" AND ")), t;
}
function g(e) {
switch (e) {
case "eq": return "equals";
case "neq": return "notequals";
case "gt": return "greaterthan";
case "gte": return "greaterthanorequal";
case "lt": return "lessthan";
case "lte": return "lessthanorequal";
case "like":
case "ilike":
case "contains": return "contains";
case "startswith": return "beginswith";
case "endswith": return "endswith";
case "in": return "in";
case "between": return "between";
case "between_inclusive": return "betweeninclusive";
case "is_null": return "empty";
case "is_not_null": return "notempty";
default: return e;
}
}
function _(e) {
return e.value === null || e.value === void 0 ? "" : Array.isArray(e.value) ? e.value.join(",") : String(e.value);
}
var v = /* @__PURE__ */ new Map();
function y(e) {
let t = a(e), n = v.get(t);
return n || (n = new b(e), v.set(t, n)), n;
}
var b = class {
constructor(e) {
this.config = {
...e,
headers: { ...e.headers }
};
}
buildUrl(e, t, n) {
let r = `${this.config.baseUrl}/${e}/${t}`;
return n && (r += `/${n}`), r;
}
baseHeaders() {
return i(this.config);
}
async fetchWithError(e, t) {
let n = await fetch(e, t), r = await n.json();
if (!n.ok) throw Error(r.error?.message || `${n.statusText} (${n.status})`);
return {
data: r,
success: !0,
error: r.error ? r.error : void 0,
metadata: {
count: n.headers.get("content-range") ? Number(n.headers.get("content-range")?.split("/")[1]) : 0,
total: n.headers.get("content-range") ? Number(n.headers.get("content-range")?.split("/")[1]) : 0,
filtered: n.headers.get("content-range") ? Number(n.headers.get("content-range")?.split("/")[1]) : 0,
offset: n.headers.get("content-range") ? Number(n.headers.get("content-range")?.split("/")[0].split("-")[0]) : 0,
limit: n.headers.get("x-limit") ? Number(n.headers.get("x-limit")) : 0
}
};
}
async read(e, t, n, i) {
let a = this.buildUrl(e, t, n), o = i ? h(i) : {};
return this.fetchWithError(a, {
method: "GET",
headers: r(this.baseHeaders(), o)
});
}
async create(e, t, n, i) {
let a = this.buildUrl(e, t), o = i ? h(i) : {};
return this.fetchWithError(a, {
method: "POST",
headers: r(this.baseHeaders(), o),
body: JSON.stringify(n)
});
}
async update(e, t, n, i, a) {
let o = this.buildUrl(e, t, n), s = a ? h(a) : {};
return this.fetchWithError(o, {
method: "PUT",
headers: r(this.baseHeaders(), s),
body: JSON.stringify(i)
});
}
async delete(e, t, n) {
let r = this.buildUrl(e, t, n);
return this.fetchWithError(r, {
method: "DELETE",
headers: this.baseHeaders()
});
}
};
//#endregion
export { b as HeaderSpecClient, c as ResolveSpecClient, d as WebSocketClient, h as buildHeaders, p as decodeHeaderValue, f as encodeHeaderValue, y as getHeaderSpecClient, s as getResolveSpecClient, u as getWebSocketClient };
+14 -12
View File
@@ -1,6 +1,6 @@
{
"name": "@warkypublic/resolvespec-js",
"version": "1.0.1",
"version": "1.0.2",
"description": "TypeScript client library for ResolveSpec REST, HeaderSpec, and WebSocket APIs",
"type": "module",
"main": "./dist/index.cjs",
@@ -38,20 +38,22 @@
"author": "Hein (Warkanum) Puth",
"license": "MIT",
"dependencies": {
"uuid": "^13.0.0"
"@warkypublic/artemis-kit": "^1.0.10",
"uuid": "^14.0.2"
},
"devDependencies": {
"@changesets/cli": "^2.29.8",
"@changesets/cli": "^3.0.3",
"@eslint/js": "^10.0.1",
"@types/jsdom": "^27.0.0",
"eslint": "^10.0.0",
"globals": "^17.3.0",
"jsdom": "^28.1.0",
"typescript": "^5.9.3",
"typescript-eslint": "^8.55.0",
"vite": "^7.3.1",
"vite-plugin-dts": "^4.5.4",
"vitest": "^4.0.18"
"@types/jsdom": "^30.0.0",
"@types/node": "^26.6.2",
"eslint": "^10.11.0",
"globals": "^17.12.0",
"jsdom": "^30.1.1",
"typescript": "^6.0.3",
"typescript-eslint": "^8.70.1",
"vite": "^8.3.0",
"vite-plugin-dts": "^5.1.1",
"vitest": "^5.0.1"
},
"engines": {
"node": ">=18"
+1283 -1293
View File
File diff suppressed because it is too large Load Diff
+5
View File
@@ -0,0 +1,5 @@
packages:
- '.'
allowBuilds:
esbuild: true
@@ -0,0 +1,65 @@
import { afterEach, describe, expect, it, vi } from 'vitest';
import { ResolveSpecClient, getResolveSpecClient } from '../resolvespec/client';
import { HeaderSpecClient, getHeaderSpecClient } from '../headerspec/client';
afterEach(() => vi.unstubAllGlobals());
for (const [name, Client, factory] of [
['ResolveSpec', ResolveSpecClient, getResolveSpecClient],
['HeaderSpec', HeaderSpecClient, getHeaderSpecClient],
] as const) {
describe(`${name} custom headers`, () => {
it('sends tenant headers on every operation and resolves collisions case-insensitively', async () => {
const fetchMock = vi.fn().mockResolvedValue({
ok: true, headers: new Headers(), json: async () => ({ success: true, data: [] }),
});
vi.stubGlobal('fetch', fetchMock);
const headers = { 'X-Tenant': 'acme', authorization: 'Basic ignored', 'content-type': 'application/custom+json', 'x-limit': '99' };
const client = new Client({ baseUrl: 'http://localhost:3000', token: 'tok', headers });
await client.read('public', 'users', undefined, { limit: 10 });
await client.create('public', 'users', {});
if (client instanceof ResolveSpecClient) {
await client.update('public', 'users', {}, '1');
await client.getMetadata('public', 'users');
} else {
await client.update('public', 'users', '1', {});
}
await client.delete('public', 'users', '1');
for (const [, init] of fetchMock.mock.calls) {
const sent = new Headers(init.headers);
expect(sent.get('x-tenant')).toBe('acme');
expect(sent.get('authorization')).toBe('Bearer tok');
expect(sent.get('content-type')).toBe('application/custom+json');
}
if (client instanceof HeaderSpecClient) {
expect(new Headers(fetchMock.mock.calls[0][1].headers).get('x-limit')).toBe('10');
}
expect(headers.authorization).toBe('Basic ignored');
expect(headers['x-limit']).toBe('99');
});
it('supports custom authentication without a token', async () => {
const fetchMock = vi.fn().mockResolvedValue({
ok: true, headers: new Headers(), json: async () => ({ success: true, data: [] }),
});
vi.stubGlobal('fetch', fetchMock);
await new Client({ baseUrl: 'http://localhost:3000', headers: { Authorization: 'Basic custom' } }).read('public', 'users');
expect(new Headers(fetchMock.mock.calls[0][1].headers).get('authorization')).toBe('Basic custom');
});
it('isolates cached clients by headers and token, and snapshots configuration', async () => {
const config = { baseUrl: 'http://tenant-cache', token: 'one', headers: { 'X-Tenant': 'acme', 'X-App': 'grid' } };
const first = factory(config);
expect(factory({ ...config, headers: { 'x-app': 'grid', 'x-tenant': 'acme' } })).toBe(first);
expect(factory({ ...config, token: 'two' })).not.toBe(first);
config.headers['X-Tenant'] = 'other';
expect(factory(config)).not.toBe(first);
const fetchMock = vi.fn().mockResolvedValue({
ok: true, headers: new Headers(), json: async () => ({ success: true, data: [] }),
});
vi.stubGlobal('fetch', fetchMock);
await first.read('public', 'users');
expect(new Headers(fetchMock.mock.calls[0][1].headers).get('x-tenant')).toBe('acme');
});
});
}
@@ -126,11 +126,22 @@ describe('encodeHeaderValue / decodeHeaderValue', () => {
expect(decoded).toBe(original);
});
it('should round-trip UTF-8 values', () => {
const original = 'café ☕ 你好';
expect(decodeHeaderValue(encodeHeaderValue(original))).toBe(original);
});
it('should decode __ prefixed values', () => {
const encoded = '__' + btoa('hello');
expect(decodeHeaderValue(encoded)).toBe('hello');
});
it('should decode UTF-8 values with the __ prefix', () => {
const bytes = new TextEncoder().encode('café ☕');
const binary = Array.from(bytes, (byte) => String.fromCharCode(byte)).join('');
expect(decodeHeaderValue('__' + btoa(binary))).toBe('café ☕');
});
it('should return plain values as-is', () => {
expect(decodeHeaderValue('plain')).toBe('plain');
});
@@ -142,6 +153,7 @@ describe('HeaderSpecClient', () => {
function mockFetch<T>(data: APIResponse<T>, ok = true) {
return vi.fn().mockResolvedValue({
ok,
headers: new Headers(),
json: () => Promise.resolve(data),
});
}
+30
View File
@@ -0,0 +1,30 @@
import type { ClientConfig } from './types';
/** Merge HTTP headers case-insensitively, preserving the winning spelling. */
export function mergeHeaders(...sources: Record<string, string>[]): Record<string, string> {
const result: Record<string, string> = {};
for (const source of sources) {
for (const [name, value] of Object.entries(source)) {
for (const existing of Object.keys(result)) {
if (existing.toLowerCase() === name.toLowerCase()) delete result[existing];
}
Object.defineProperty(result, name, { value, enumerable: true, configurable: true, writable: true });
}
}
return result;
}
export function clientHeaders(config: ClientConfig): Record<string, string> {
return mergeHeaders(
{ 'Content-Type': 'application/json' },
config.headers ?? {},
config.token ? { Authorization: `Bearer ${config.token}` } : {},
);
}
export function clientCacheKey(config: ClientConfig): string {
const headers = Object.entries(clientHeaders(config))
.map(([name, value]) => [name.toLowerCase(), value])
.sort(([a], [b]) => a.localeCompare(b));
return JSON.stringify([config.baseUrl, headers]);
}
+2
View File
@@ -126,4 +126,6 @@ export interface TableMetadata {
export interface ClientConfig {
baseUrl: string;
token?: string;
/** Custom HTTP headers. Token and HeaderSpec query options take precedence. */
headers?: Record<string, string>;
}
+10 -20
View File
@@ -1,3 +1,5 @@
import { clientCacheKey, clientHeaders, mergeHeaders } from '../common/http';
import { b64DecodeUnicode, b64EncodeUnicode } from '@warkypublic/artemis-kit/base64';
import type {
APIResponse,
ClientConfig,
@@ -12,10 +14,7 @@ import type {
* Encode a value with base64 and ZIP_ prefix for complex header values.
*/
export function encodeHeaderValue(value: string): string {
if (typeof btoa === "function") {
return "ZIP_" + btoa(value);
}
return "ZIP_" + Buffer.from(value, "utf-8").toString("base64");
return "ZIP_" + b64EncodeUnicode(value);
}
/**
@@ -41,10 +40,7 @@ export function decodeHeaderValue(value: string): string {
}
function decodeBase64(str: string): string {
if (typeof atob === "function") {
return atob(str);
}
return Buffer.from(str, "base64").toString("utf-8");
return b64DecodeUnicode(str);
}
/**
@@ -203,7 +199,7 @@ function formatFilterValue(filter: FilterOption): string {
const instances = new Map<string, HeaderSpecClient>();
export function getHeaderSpecClient(config: ClientConfig): HeaderSpecClient {
const key = config.baseUrl;
const key = clientCacheKey(config);
let instance = instances.get(key);
if (!instance) {
instance = new HeaderSpecClient(config);
@@ -222,7 +218,7 @@ export class HeaderSpecClient {
private config: ClientConfig;
constructor(config: ClientConfig) {
this.config = config;
this.config = { ...config, headers: { ...config.headers } };
}
private buildUrl(schema: string, entity: string, id?: string): string {
@@ -234,13 +230,7 @@ export class HeaderSpecClient {
}
private baseHeaders(): Record<string, string> {
const headers: Record<string, string> = {
"Content-Type": "application/json",
};
if (this.config.token) {
headers["Authorization"] = `Bearer ${this.config.token}`;
}
return headers;
return clientHeaders(this.config);
}
private async fetchWithError<T>(
@@ -296,7 +286,7 @@ export class HeaderSpecClient {
const optHeaders = options ? buildHeaders(options) : {};
return this.fetchWithError<T>(url, {
method: "GET",
headers: { ...this.baseHeaders(), ...optHeaders },
headers: mergeHeaders(this.baseHeaders(), optHeaders),
});
}
@@ -310,7 +300,7 @@ export class HeaderSpecClient {
const optHeaders = options ? buildHeaders(options) : {};
return this.fetchWithError<T>(url, {
method: "POST",
headers: { ...this.baseHeaders(), ...optHeaders },
headers: mergeHeaders(this.baseHeaders(), optHeaders),
body: JSON.stringify(data),
});
}
@@ -326,7 +316,7 @@ export class HeaderSpecClient {
const optHeaders = options ? buildHeaders(options) : {};
return this.fetchWithError<T>(url, {
method: "PUT",
headers: { ...this.baseHeaders(), ...optHeaders },
headers: mergeHeaders(this.baseHeaders(), optHeaders),
body: JSON.stringify(data),
});
}
+4 -11
View File
@@ -1,9 +1,10 @@
import { clientCacheKey, clientHeaders } from '../common/http';
import type { ClientConfig, APIResponse, TableMetadata, Options, RequestBody } from '../common/types';
const instances = new Map<string, ResolveSpecClient>();
export function getResolveSpecClient(config: ClientConfig): ResolveSpecClient {
const key = config.baseUrl;
const key = clientCacheKey(config);
let instance = instances.get(key);
if (!instance) {
instance = new ResolveSpecClient(config);
@@ -16,7 +17,7 @@ export class ResolveSpecClient {
private config: ClientConfig;
constructor(config: ClientConfig) {
this.config = config;
this.config = { ...config, headers: { ...config.headers } };
}
private buildUrl(schema: string, entity: string, id?: string): string {
@@ -28,15 +29,7 @@ export class ResolveSpecClient {
}
private baseHeaders(): HeadersInit {
const headers: Record<string, string> = {
'Content-Type': 'application/json',
};
if (this.config.token) {
headers['Authorization'] = `Bearer ${this.config.token}`;
}
return headers;
return clientHeaders(this.config);
}
private async fetchWithError<T>(url: string, options: RequestInit): Promise<APIResponse<T>> {
+1 -1
View File
@@ -14,7 +14,7 @@ export default defineConfig({
fileName: (format) => `index.${format === 'es' ? 'js' : 'cjs'}`,
},
rollupOptions: {
external: ['uuid', 'semver'],
external: ['uuid', 'semver', '@warkypublic/artemis-kit/base64'],
},
},
});