mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-09-28 19:12:00 +00:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
6687a7a5cd | ||
|
|
e8fbbede7e | ||
|
|
7f8982fa35 | ||
|
|
b587cbd3c4 | ||
|
|
a220338eea | ||
|
|
20c67166d0 | ||
|
|
749dad4ed1 | ||
|
|
d6c5740f9c | ||
|
|
817b781c88 | ||
|
|
87eaa9e18c | ||
|
|
4f6878099b | ||
|
|
0d8b136b91 | ||
|
|
6de9be0ae7 | ||
|
|
82e923b16e | ||
|
|
9a664593f0 | ||
|
|
6e3124e4e0 | ||
|
|
d5de48011b | ||
|
|
6bd6a6f164 | ||
|
|
1885ce016b | ||
|
|
eeb7ba04d8 | ||
|
|
e957753ce4 | ||
|
|
206edd4bfd | ||
|
|
f841d58c59 |
@@ -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)
|
||||
}
|
||||
@@ -0,0 +1,402 @@
|
||||
package common
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"regexp"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// This file implements a single canonical parser + SQL builder for column
|
||||
// references that traverse into JSON / JSONB values. It is used by SELECT,
|
||||
// WHERE (filter) and ORDER BY handling so that all three treat JSON access
|
||||
// consistently and safely.
|
||||
//
|
||||
// Supported input syntaxes (all PostgreSQL-oriented):
|
||||
//
|
||||
// data->>'city' arrow chain, text extraction
|
||||
// data->'addr'->>'city' nested arrow chain
|
||||
// data->2->>'name' arrow chain with array index
|
||||
// data#>>'{addr,city}' hash-path, text extraction
|
||||
// data#>'{addr,city}' hash-path, jsonb result
|
||||
// data.addr.city dotted shorthand (Ambiguous: caller must
|
||||
// confirm "data" is a JSON column)
|
||||
// data->>'age'::int trailing cast (whitelisted targets only)
|
||||
// (data->>'city') AS city parenthesised, with output alias
|
||||
//
|
||||
// JSON path segments are never interpolated into SQL: SQL() emits a `#>>` /
|
||||
// `#>` operator with the path bound as a single `text[]` parameter.
|
||||
|
||||
// ColumnRef is a parsed reference to a (possibly JSON-traversing) column.
|
||||
type ColumnRef struct {
|
||||
// Base is the bare base column name, e.g. "data". Always a simple
|
||||
// identifier ([A-Za-z_][A-Za-z0-9_]*); qualified names are rejected.
|
||||
Base string
|
||||
// Path is the JSON key / array-index path, e.g. ["address", "city"].
|
||||
// Empty for a plain column reference.
|
||||
Path []string
|
||||
// AsText is true when the final extraction should yield text (->> / #>>)
|
||||
// rather than jsonb (-> / #>).
|
||||
AsText bool
|
||||
// Cast is a normalised SQL type name to cast the whole expression to
|
||||
// (e.g. "integer", "numeric", "timestamptz"), or "" for no cast.
|
||||
Cast string
|
||||
// Alias is a validated output identifier for `AS <alias>`, or "".
|
||||
Alias string
|
||||
// Ambiguous is true when Path was produced from the dotted "a.b.c"
|
||||
// shorthand. The caller MUST verify that Base is a JSON column
|
||||
// (reflection.IsJSONColumn) before treating this as a JSON expression,
|
||||
// otherwise "a.b" is an ordinary table-qualified column.
|
||||
Ambiguous bool
|
||||
}
|
||||
|
||||
var (
|
||||
reSimpleIdent = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*$`)
|
||||
reSimpleSegment = regexp.MustCompile(`^[A-Za-z0-9_]+$`)
|
||||
reAliasSuffix = regexp.MustCompile(`(?i)\s+AS\s+("?[A-Za-z_][A-Za-z0-9_]*"?)\s*$`)
|
||||
reArrowStep = regexp.MustCompile(`^\s*(->>|->)\s*(?:'((?:[^']|'')*)'|(\d+))\s*`)
|
||||
)
|
||||
|
||||
const (
|
||||
maxJSONPathDepth = 32
|
||||
maxJSONSegmentSize = 128
|
||||
)
|
||||
|
||||
// castAliases maps accepted cast spellings to their canonical PostgreSQL type.
|
||||
var castAliases = map[string]string{
|
||||
"int": "integer",
|
||||
"int4": "integer",
|
||||
"integer": "integer",
|
||||
"int2": "smallint",
|
||||
"smallint": "smallint",
|
||||
"int8": "bigint",
|
||||
"bigint": "bigint",
|
||||
"numeric": "numeric",
|
||||
"decimal": "numeric",
|
||||
"real": "real",
|
||||
"float4": "real",
|
||||
"float": "double precision",
|
||||
"float8": "double precision",
|
||||
"double precision": "double precision",
|
||||
"bool": "boolean",
|
||||
"boolean": "boolean",
|
||||
"text": "text",
|
||||
"varchar": "text",
|
||||
"uuid": "uuid",
|
||||
"date": "date",
|
||||
"time": "time",
|
||||
"timestamp": "timestamp",
|
||||
"timestamptz": "timestamptz",
|
||||
"json": "json",
|
||||
"jsonb": "jsonb",
|
||||
}
|
||||
|
||||
// NormalizeCastTarget returns the canonical PostgreSQL type name for a
|
||||
// user-supplied cast spelling, and whether it is on the allowlist.
|
||||
func NormalizeCastTarget(s string) (string, bool) {
|
||||
c, ok := castAliases[strings.ToLower(strings.TrimSpace(s))]
|
||||
return c, ok
|
||||
}
|
||||
|
||||
// ParseColumnRef parses a column token that traverses into a JSON value.
|
||||
//
|
||||
// ok is true only when the token carries JSON traversal syntax (arrow chain,
|
||||
// hash-path, or dotted shorthand with at least one sub-key). For a plain
|
||||
// column name — with or without an alias/cast — ok is false and the caller
|
||||
// should handle the token the way it did before.
|
||||
//
|
||||
// When ok is true and ref.Ambiguous is true, the caller must confirm that
|
||||
// ref.Base is a JSON column before using ref.SQL; otherwise the dotted token
|
||||
// is an ordinary "table.column" reference.
|
||||
func ParseColumnRef(raw string) (ColumnRef, bool) {
|
||||
expr := strings.TrimSpace(raw)
|
||||
if expr == "" {
|
||||
return ColumnRef{}, false
|
||||
}
|
||||
|
||||
var ref ColumnRef
|
||||
|
||||
// 1. Trailing `AS <alias>`.
|
||||
if m := reAliasSuffix.FindStringSubmatch(expr); m != nil {
|
||||
ref.Alias = strings.Trim(m[1], `"`)
|
||||
expr = strings.TrimSpace(expr[:len(expr)-len(m[0])])
|
||||
if expr == "" {
|
||||
return ColumnRef{}, false
|
||||
}
|
||||
}
|
||||
|
||||
// 2. Trailing `::<type>` cast (take the last `::` in the string).
|
||||
if idx := strings.LastIndex(expr, "::"); idx != -1 {
|
||||
candidate := strings.TrimSpace(expr[idx+2:])
|
||||
if canonical, allowed := NormalizeCastTarget(candidate); allowed {
|
||||
ref.Cast = canonical
|
||||
expr = strings.TrimSpace(expr[:idx])
|
||||
} else if candidate != "" && looksLikeCastTail(candidate) {
|
||||
// An explicit but unsupported cast target — reject rather than
|
||||
// silently dropping it.
|
||||
return ColumnRef{}, false
|
||||
}
|
||||
}
|
||||
|
||||
// 3. One layer of wrapping parentheses: "(expr)" -> "expr".
|
||||
if wrapped, ok := stripWrappingParens(expr); ok {
|
||||
expr = strings.TrimSpace(wrapped)
|
||||
if expr == "" {
|
||||
return ColumnRef{}, false
|
||||
}
|
||||
}
|
||||
|
||||
// 4. Parse the core expression.
|
||||
switch {
|
||||
case strings.Contains(expr, "#>>") || strings.Contains(expr, "#>"):
|
||||
if !parseHashPath(expr, &ref) {
|
||||
return ColumnRef{}, false
|
||||
}
|
||||
case strings.Contains(expr, "->"):
|
||||
if !parseArrowChain(expr, &ref) {
|
||||
return ColumnRef{}, false
|
||||
}
|
||||
case strings.Contains(expr, "."):
|
||||
if !parseDottedPath(expr, &ref) {
|
||||
return ColumnRef{}, false
|
||||
}
|
||||
default:
|
||||
// Plain column — nothing JSON about it.
|
||||
return ColumnRef{}, false
|
||||
}
|
||||
|
||||
if !validateRef(&ref) {
|
||||
return ColumnRef{}, false
|
||||
}
|
||||
return ref, true
|
||||
}
|
||||
|
||||
// SQL renders the reference as a parameterised SQL expression plus its args.
|
||||
// tableAlias, when non-empty, qualifies the base column (each dot-separated
|
||||
// part is quoted independently, so "public.users" -> `"public"."users"`).
|
||||
func (r ColumnRef) SQL(tableAlias string) (expr string, args []interface{}) {
|
||||
base := quoteQualifiedIdent(r.Base)
|
||||
if tableAlias != "" {
|
||||
base = quoteQualifiedIdent(tableAlias) + "." + QuoteIdent(r.Base)
|
||||
}
|
||||
|
||||
if len(r.Path) == 0 {
|
||||
if r.Cast != "" {
|
||||
return fmt.Sprintf("(%s)::%s", base, r.Cast), nil
|
||||
}
|
||||
return base, nil
|
||||
}
|
||||
|
||||
op := "#>"
|
||||
if r.AsText {
|
||||
op = "#>>"
|
||||
}
|
||||
expr = fmt.Sprintf("(%s %s ?::text[])", base, op)
|
||||
args = []interface{}{pgTextArrayLiteral(r.Path)}
|
||||
|
||||
if r.Cast != "" {
|
||||
expr = fmt.Sprintf("(%s)::%s", expr, r.Cast)
|
||||
}
|
||||
return expr, args
|
||||
}
|
||||
|
||||
// OutputAlias returns the alias to use for this reference in a SELECT list:
|
||||
// the explicit alias when given, otherwise a deterministic name derived from
|
||||
// the base column and path (e.g. "data_address_city").
|
||||
func (r ColumnRef) OutputAlias() string {
|
||||
if r.Alias != "" {
|
||||
return r.Alias
|
||||
}
|
||||
if len(r.Path) == 0 {
|
||||
return r.Base
|
||||
}
|
||||
parts := make([]string, 0, len(r.Path)+1)
|
||||
parts = append(parts, r.Base)
|
||||
for _, p := range r.Path {
|
||||
parts = append(parts, sanitizeAliasPart(p))
|
||||
}
|
||||
return strings.Join(parts, "_")
|
||||
}
|
||||
|
||||
// ── parsing helpers ─────────────────────────────────────────────────────────
|
||||
|
||||
func parseHashPath(expr string, ref *ColumnRef) bool {
|
||||
op := "#>>"
|
||||
ref.AsText = true
|
||||
if !strings.Contains(expr, "#>>") {
|
||||
op = "#>"
|
||||
ref.AsText = false
|
||||
}
|
||||
parts := strings.SplitN(expr, op, 2)
|
||||
if len(parts) != 2 {
|
||||
return false
|
||||
}
|
||||
ref.Base = strings.TrimSpace(parts[0])
|
||||
|
||||
rhs := strings.TrimSpace(parts[1])
|
||||
// Expect a single-quoted array literal: '{a,b,c}'
|
||||
if len(rhs) < 2 || rhs[0] != '\'' || rhs[len(rhs)-1] != '\'' {
|
||||
return false
|
||||
}
|
||||
rhs = rhs[1 : len(rhs)-1]
|
||||
rhs = strings.TrimSpace(rhs)
|
||||
rhs = strings.TrimPrefix(rhs, "{")
|
||||
rhs = strings.TrimSuffix(rhs, "}")
|
||||
if strings.TrimSpace(rhs) == "" {
|
||||
return false
|
||||
}
|
||||
for _, seg := range strings.Split(rhs, ",") {
|
||||
seg = strings.TrimSpace(seg)
|
||||
seg = strings.Trim(seg, `"`)
|
||||
if seg == "" {
|
||||
return false
|
||||
}
|
||||
ref.Path = append(ref.Path, seg)
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func parseArrowChain(expr string, ref *ColumnRef) bool {
|
||||
arrowIdx := strings.Index(expr, "->")
|
||||
if arrowIdx <= 0 {
|
||||
return false
|
||||
}
|
||||
ref.Base = strings.TrimSpace(expr[:arrowIdx])
|
||||
|
||||
rest := expr[arrowIdx:]
|
||||
for strings.TrimSpace(rest) != "" {
|
||||
m := reArrowStep.FindStringSubmatch(rest)
|
||||
if m == nil {
|
||||
return false
|
||||
}
|
||||
ref.AsText = m[1] == "->>"
|
||||
if m[3] != "" {
|
||||
// unquoted array index
|
||||
ref.Path = append(ref.Path, m[3])
|
||||
} else {
|
||||
// quoted key; unescape doubled single quotes
|
||||
ref.Path = append(ref.Path, strings.ReplaceAll(m[2], "''", "'"))
|
||||
}
|
||||
rest = rest[len(m[0]):]
|
||||
}
|
||||
return len(ref.Path) > 0
|
||||
}
|
||||
|
||||
func parseDottedPath(expr string, ref *ColumnRef) bool {
|
||||
segs := strings.Split(expr, ".")
|
||||
if len(segs) < 2 {
|
||||
return false
|
||||
}
|
||||
for i, s := range segs {
|
||||
s = strings.TrimSpace(s)
|
||||
if !reSimpleSegment.MatchString(s) {
|
||||
return false
|
||||
}
|
||||
if i == 0 {
|
||||
ref.Base = s
|
||||
} else {
|
||||
ref.Path = append(ref.Path, s)
|
||||
}
|
||||
}
|
||||
ref.AsText = true
|
||||
ref.Ambiguous = true
|
||||
return true
|
||||
}
|
||||
|
||||
func validateRef(ref *ColumnRef) bool {
|
||||
if !reSimpleIdent.MatchString(ref.Base) {
|
||||
return false
|
||||
}
|
||||
if len(ref.Path) == 0 || len(ref.Path) > maxJSONPathDepth {
|
||||
return false
|
||||
}
|
||||
for _, seg := range ref.Path {
|
||||
if seg == "" || len(seg) > maxJSONSegmentSize || strings.ContainsRune(seg, 0) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
if ref.Alias != "" && !reSimpleIdent.MatchString(ref.Alias) {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// looksLikeCastTail reports whether s is plausibly meant as a `::type` target
|
||||
// (letters/digits/spaces only) rather than, say, part of a JSON operator.
|
||||
func looksLikeCastTail(s string) bool {
|
||||
for _, r := range s {
|
||||
isCastChar := r == ' ' || r == '_' ||
|
||||
(r >= 'a' && r <= 'z') || (r >= 'A' && r <= 'Z') || (r >= '0' && r <= '9')
|
||||
if !isCastChar {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// stripWrappingParens removes one layer of parentheses when they wrap the whole
|
||||
// expression, e.g. "(a->>'b')" -> "a->>'b'". It respects single-quoted strings.
|
||||
func stripWrappingParens(s string) (string, bool) {
|
||||
s = strings.TrimSpace(s)
|
||||
if len(s) < 2 || s[0] != '(' || s[len(s)-1] != ')' {
|
||||
return s, false
|
||||
}
|
||||
depth := 0
|
||||
inQuote := false
|
||||
for i := 0; i < len(s); i++ {
|
||||
c := s[i]
|
||||
switch {
|
||||
case c == '\'':
|
||||
inQuote = !inQuote
|
||||
case inQuote:
|
||||
// skip
|
||||
case c == '(':
|
||||
depth++
|
||||
case c == ')':
|
||||
depth--
|
||||
if depth == 0 && i != len(s)-1 {
|
||||
// closing paren is not the last char -> not a full wrap
|
||||
return s, false
|
||||
}
|
||||
}
|
||||
}
|
||||
if depth != 0 {
|
||||
return s, false
|
||||
}
|
||||
return s[1 : len(s)-1], true
|
||||
}
|
||||
|
||||
// pgTextArrayLiteral builds a PostgreSQL text[] array literal ("{a,b,c}") from
|
||||
// path segments, quoting and escaping any segment that is not a bare word.
|
||||
func pgTextArrayLiteral(segs []string) string {
|
||||
escaper := strings.NewReplacer(`\`, `\\`, `"`, `\"`)
|
||||
parts := make([]string, len(segs))
|
||||
for i, s := range segs {
|
||||
if reSimpleSegment.MatchString(s) {
|
||||
parts[i] = s
|
||||
} else {
|
||||
parts[i] = `"` + escaper.Replace(s) + `"`
|
||||
}
|
||||
}
|
||||
return "{" + strings.Join(parts, ",") + "}"
|
||||
}
|
||||
|
||||
// quoteQualifiedIdent quotes each dot-separated part of an identifier.
|
||||
func quoteQualifiedIdent(ident string) string {
|
||||
parts := strings.Split(ident, ".")
|
||||
for i, p := range parts {
|
||||
parts[i] = QuoteIdent(p)
|
||||
}
|
||||
return strings.Join(parts, ".")
|
||||
}
|
||||
|
||||
func sanitizeAliasPart(s string) string {
|
||||
var b strings.Builder
|
||||
for _, r := range s {
|
||||
if (r >= 'a' && r <= 'z') || (r >= 'A' && r <= 'Z') || (r >= '0' && r <= '9') || r == '_' {
|
||||
b.WriteRune(r)
|
||||
} else {
|
||||
b.WriteRune('_')
|
||||
}
|
||||
}
|
||||
return b.String()
|
||||
}
|
||||
@@ -0,0 +1,274 @@
|
||||
package common
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestParseColumnRef_Valid(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
input string
|
||||
base string
|
||||
path []string
|
||||
asText bool
|
||||
cast string
|
||||
alias string
|
||||
ambiguous bool
|
||||
}{
|
||||
{
|
||||
name: "arrow text extraction",
|
||||
input: "data->>'city'",
|
||||
base: "data", path: []string{"city"}, asText: true,
|
||||
},
|
||||
{
|
||||
name: "arrow whitespace tolerant",
|
||||
input: "data ->> 'city'",
|
||||
base: "data", path: []string{"city"}, asText: true,
|
||||
},
|
||||
{
|
||||
name: "nested arrow chain",
|
||||
input: "data->'address'->>'city'",
|
||||
base: "data", path: []string{"address", "city"}, asText: true,
|
||||
},
|
||||
{
|
||||
name: "arrow jsonb result",
|
||||
input: "data->'address'",
|
||||
base: "data", path: []string{"address"}, asText: false,
|
||||
},
|
||||
{
|
||||
name: "arrow array index",
|
||||
input: "items->0->>'name'",
|
||||
base: "items", path: []string{"0", "name"}, asText: true,
|
||||
},
|
||||
{
|
||||
name: "hash path text",
|
||||
input: "data#>>'{address,city}'",
|
||||
base: "data", path: []string{"address", "city"}, asText: true,
|
||||
},
|
||||
{
|
||||
name: "hash path jsonb",
|
||||
input: "data#>'{address,city}'",
|
||||
base: "data", path: []string{"address", "city"}, asText: false,
|
||||
},
|
||||
{
|
||||
name: "dotted shorthand",
|
||||
input: "data.address.city",
|
||||
base: "data", path: []string{"address", "city"}, asText: true, ambiguous: true,
|
||||
},
|
||||
{
|
||||
name: "trailing cast",
|
||||
input: "data->>'age'::int",
|
||||
base: "data", path: []string{"age"}, asText: true, cast: "integer",
|
||||
},
|
||||
{
|
||||
name: "cast normalises",
|
||||
input: "data->>'ts'::timestamptz",
|
||||
base: "data", path: []string{"ts"}, asText: true, cast: "timestamptz",
|
||||
},
|
||||
{
|
||||
name: "parenthesised with alias",
|
||||
input: "(data->>'city') AS city_name",
|
||||
base: "data", path: []string{"city"}, asText: true, alias: "city_name",
|
||||
},
|
||||
{
|
||||
name: "paren wrap and cast",
|
||||
input: "(data->>'age')::numeric",
|
||||
base: "data", path: []string{"age"}, asText: true, cast: "numeric",
|
||||
},
|
||||
{
|
||||
name: "quoted key with spaces",
|
||||
input: "data->>'key with space'",
|
||||
base: "data", path: []string{"key with space"}, asText: true,
|
||||
},
|
||||
{
|
||||
name: "quoted key with escaped quote",
|
||||
input: "data->>'o''brien'",
|
||||
base: "data", path: []string{"o'brien"}, asText: true,
|
||||
},
|
||||
{
|
||||
name: "relation column is ambiguous json",
|
||||
input: "orders.total",
|
||||
base: "orders", path: []string{"total"}, asText: true, ambiguous: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
ref, ok := ParseColumnRef(tc.input)
|
||||
if !ok {
|
||||
t.Fatalf("ParseColumnRef(%q) returned ok=false", tc.input)
|
||||
}
|
||||
if ref.Base != tc.base {
|
||||
t.Errorf("Base = %q, want %q", ref.Base, tc.base)
|
||||
}
|
||||
if !reflect.DeepEqual(ref.Path, tc.path) {
|
||||
t.Errorf("Path = %#v, want %#v", ref.Path, tc.path)
|
||||
}
|
||||
if ref.AsText != tc.asText {
|
||||
t.Errorf("AsText = %v, want %v", ref.AsText, tc.asText)
|
||||
}
|
||||
if ref.Cast != tc.cast {
|
||||
t.Errorf("Cast = %q, want %q", ref.Cast, tc.cast)
|
||||
}
|
||||
if ref.Alias != tc.alias {
|
||||
t.Errorf("Alias = %q, want %q", ref.Alias, tc.alias)
|
||||
}
|
||||
if ref.Ambiguous != tc.ambiguous {
|
||||
t.Errorf("Ambiguous = %v, want %v", ref.Ambiguous, tc.ambiguous)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseColumnRef_NotJSON(t *testing.T) {
|
||||
// These must return ok=false so callers fall back to their normal handling.
|
||||
inputs := []string{
|
||||
"",
|
||||
" ",
|
||||
"name",
|
||||
"data",
|
||||
"created_at",
|
||||
"(id)",
|
||||
}
|
||||
for _, in := range inputs {
|
||||
if ref, ok := ParseColumnRef(in); ok {
|
||||
t.Errorf("ParseColumnRef(%q) = %+v, ok=true; want ok=false", in, ref)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseColumnRef_Rejected(t *testing.T) {
|
||||
// Malformed or unsafe tokens must be rejected outright.
|
||||
inputs := []string{
|
||||
"data->>'x'::bogus", // cast not on allowlist
|
||||
"data->>'x' AS 1bad", // invalid alias
|
||||
"(data->>'a') OR (x->>'b')", // not a single wrapped expr
|
||||
"data->>'x'); DROP TABLE users; --", // injection attempt
|
||||
"data->b", // unquoted non-numeric key
|
||||
"data->>''", // empty key
|
||||
"data#>>'{}'", // empty hash path
|
||||
"data#>>address", // hash path not a quoted literal
|
||||
"weird col->>'x'", // base not an identifier
|
||||
"data.address.city.but.way.too...deep.", // trailing dot -> empty segment
|
||||
}
|
||||
for _, in := range inputs {
|
||||
if ref, ok := ParseColumnRef(in); ok {
|
||||
t.Errorf("ParseColumnRef(%q) = %+v, ok=true; want rejected", in, ref)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestColumnRef_SQL(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
ref ColumnRef
|
||||
alias string
|
||||
wantExpr string
|
||||
wantArgs []interface{}
|
||||
}{
|
||||
{
|
||||
name: "text extraction qualified",
|
||||
ref: ColumnRef{Base: "data", Path: []string{"address", "city"}, AsText: true},
|
||||
alias: "u",
|
||||
wantExpr: `("u"."data" #>> ?::text[])`,
|
||||
wantArgs: []interface{}{"{address,city}"},
|
||||
},
|
||||
{
|
||||
name: "jsonb extraction unqualified",
|
||||
ref: ColumnRef{Base: "data", Path: []string{"a"}, AsText: false},
|
||||
alias: "",
|
||||
wantExpr: `("data" #> ?::text[])`,
|
||||
wantArgs: []interface{}{"{a}"},
|
||||
},
|
||||
{
|
||||
name: "with cast",
|
||||
ref: ColumnRef{Base: "data", Path: []string{"age"}, AsText: true, Cast: "integer"},
|
||||
alias: "t",
|
||||
wantExpr: `(("t"."data" #>> ?::text[]))::integer`,
|
||||
wantArgs: []interface{}{"{age}"},
|
||||
},
|
||||
{
|
||||
name: "schema qualified alias",
|
||||
ref: ColumnRef{Base: "data", Path: []string{"k"}, AsText: true},
|
||||
alias: "public.users",
|
||||
wantExpr: `("public"."users"."data" #>> ?::text[])`,
|
||||
wantArgs: []interface{}{"{k}"},
|
||||
},
|
||||
{
|
||||
name: "key needing quoting",
|
||||
ref: ColumnRef{Base: "data", Path: []string{"key with space", `ev"il`}, AsText: true},
|
||||
alias: "",
|
||||
wantExpr: `("data" #>> ?::text[])`,
|
||||
wantArgs: []interface{}{`{"key with space","ev\"il"}`},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
expr, args := tc.ref.SQL(tc.alias)
|
||||
if expr != tc.wantExpr {
|
||||
t.Errorf("expr = %q, want %q", expr, tc.wantExpr)
|
||||
}
|
||||
if !reflect.DeepEqual(args, tc.wantArgs) {
|
||||
t.Errorf("args = %#v, want %#v", args, tc.wantArgs)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestColumnRef_SQL_RoundTrip(t *testing.T) {
|
||||
ref, ok := ParseColumnRef("profile->'contact'->>'email'")
|
||||
if !ok {
|
||||
t.Fatal("parse failed")
|
||||
}
|
||||
expr, args := ref.SQL("customers")
|
||||
wantExpr := `("customers"."profile" #>> ?::text[])`
|
||||
if expr != wantExpr {
|
||||
t.Errorf("expr = %q, want %q", expr, wantExpr)
|
||||
}
|
||||
if len(args) != 1 || args[0] != "{contact,email}" {
|
||||
t.Errorf("args = %#v, want [{contact,email}]", args)
|
||||
}
|
||||
}
|
||||
|
||||
func TestColumnRef_OutputAlias(t *testing.T) {
|
||||
cases := []struct {
|
||||
ref ColumnRef
|
||||
want string
|
||||
}{
|
||||
{ColumnRef{Base: "data", Path: []string{"address", "city"}, AsText: true}, "data_address_city"},
|
||||
{ColumnRef{Base: "data", Path: []string{"city"}, Alias: "city"}, "city"},
|
||||
{ColumnRef{Base: "data", Path: []string{"weird key"}}, "data_weird_key"},
|
||||
{ColumnRef{Base: "data"}, "data"},
|
||||
}
|
||||
for _, c := range cases {
|
||||
if got := c.ref.OutputAlias(); got != c.want {
|
||||
t.Errorf("OutputAlias(%+v) = %q, want %q", c.ref, got, c.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeCastTarget(t *testing.T) {
|
||||
ok := map[string]string{
|
||||
"int": "integer",
|
||||
"INT": "integer",
|
||||
" bigint ": "bigint",
|
||||
"decimal": "numeric",
|
||||
"float8": "double precision",
|
||||
"bool": "boolean",
|
||||
"timestamptz": "timestamptz",
|
||||
"uuid": "uuid",
|
||||
}
|
||||
for in, want := range ok {
|
||||
got, allowed := NormalizeCastTarget(in)
|
||||
if !allowed || got != want {
|
||||
t.Errorf("NormalizeCastTarget(%q) = %q, %v; want %q, true", in, got, allowed, want)
|
||||
}
|
||||
}
|
||||
for _, in := range []string{"", "regclass", "int; drop", "text[]"} {
|
||||
if got, allowed := NormalizeCastTarget(in); allowed {
|
||||
t.Errorf("NormalizeCastTarget(%q) = %q, true; want not allowed", in, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,209 @@
|
||||
package common
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
||||
)
|
||||
|
||||
// This file wires the canonical JSON column parser (json_column.go) into the
|
||||
// three query-building paths that every spec handler shares: SELECT column
|
||||
// lists, WHERE filters and ORDER BY. The helpers here are the single place
|
||||
// those paths call so that JSON access is resolved (and made injection-safe)
|
||||
// identically everywhere. They mirror the style of BuildSpatialCondition /
|
||||
// BuildVectorCondition: a boolean ok result tells the caller whether the token
|
||||
// was a JSON reference it should take over, otherwise the caller keeps its
|
||||
// existing (non-JSON) behaviour.
|
||||
|
||||
// jsonComparisonOps are the operators for which a JSON text extraction should be
|
||||
// cast to a concrete type when the value looks numeric — otherwise "10" < "9".
|
||||
var jsonComparisonOps = map[string]bool{
|
||||
"gt": true, "greater_than": true, ">": true,
|
||||
"gte": true, "greater_than_equals": true, "ge": true, ">=": true,
|
||||
"lt": true, "less_than": true, "<": true,
|
||||
"lte": true, "less_than_equals": true, "le": true, "<=": true,
|
||||
"between": true, "between_inclusive": true,
|
||||
}
|
||||
|
||||
// ResolveJSONColumnRef parses token and, when it is a usable JSON reference for
|
||||
// model, returns the parsed ColumnRef. For the dotted "a.b" shorthand (which is
|
||||
// otherwise indistinguishable from a table-qualified column) ok is true only
|
||||
// when model confirms the base is a JSON column.
|
||||
func ResolveJSONColumnRef(model interface{}, token string) (ColumnRef, bool) {
|
||||
ref, ok := ParseColumnRef(token)
|
||||
if !ok {
|
||||
return ColumnRef{}, false
|
||||
}
|
||||
if ref.Ambiguous && !reflection.IsJSONColumn(model, ref.Base) {
|
||||
return ColumnRef{}, false
|
||||
}
|
||||
return ref, true
|
||||
}
|
||||
|
||||
// IsJSONColumnToken reports whether token is a JSON reference this package can
|
||||
// resolve for model (arrow/hash syntax always; dotted shorthand only when the
|
||||
// base is a JSON column).
|
||||
func IsJSONColumnToken(model interface{}, token string) bool {
|
||||
_, ok := ResolveJSONColumnRef(model, token)
|
||||
return ok
|
||||
}
|
||||
|
||||
// ResolveJSONColumnExpr resolves a raw column token that traverses into a JSON
|
||||
// value into a parameterised SQL expression plus its args and a deterministic
|
||||
// output alias. ok is false when the token is not a JSON reference, in which
|
||||
// case the caller should handle it the way it did before.
|
||||
//
|
||||
// tableAlias, when non-empty, qualifies the base column.
|
||||
func ResolveJSONColumnExpr(model interface{}, tableAlias, token string) (expr string, args []interface{}, alias string, ok bool) {
|
||||
ref, ok := ResolveJSONColumnRef(model, token)
|
||||
if !ok {
|
||||
return "", nil, "", false
|
||||
}
|
||||
expr, args = ref.SQL(tableAlias)
|
||||
return expr, args, ref.OutputAlias(), true
|
||||
}
|
||||
|
||||
// ApplySelectColumns adds the requested columns to query, resolving any that are
|
||||
// JSON sub-field references (data->>'x', data#>>'{a,b}', or the dotted data.x
|
||||
// shorthand for a JSON column) into safe parameterised expressions with a
|
||||
// deterministic alias. Plain columns are passed through reflection.ExtractSourceColumn
|
||||
// exactly as before. tableAlias, when non-empty, qualifies JSON base columns.
|
||||
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 {
|
||||
query = query.ColumnExpr(expr+" AS "+QuoteIdent(alias), args...)
|
||||
continue
|
||||
}
|
||||
query = query.Column(reflection.ExtractSourceColumn(col))
|
||||
}
|
||||
return query
|
||||
}
|
||||
|
||||
// BuildJSONFilterCondition builds a complete WHERE condition for a JSON column
|
||||
// token. ok is false when the token is not a JSON reference or the operator is
|
||||
// not one this builder handles (the caller then keeps its existing behaviour).
|
||||
//
|
||||
// The JSON path is always bound as a parameter, never interpolated. When the
|
||||
// reference carries no explicit ::cast and the operator is an ordered
|
||||
// comparison against a numeric value, the extracted text is cast to numeric so
|
||||
// the comparison is numeric rather than lexical.
|
||||
func BuildJSONFilterCondition(model interface{}, tableAlias, token, operator string, value interface{}) (condition string, args []interface{}, ok bool) {
|
||||
ref, ok := ResolveJSONColumnRef(model, token)
|
||||
if !ok {
|
||||
return "", nil, false
|
||||
}
|
||||
|
||||
op := strings.ToLower(strings.TrimSpace(operator))
|
||||
|
||||
// Infer a cast for ordered comparisons on numeric values so "10" > "9".
|
||||
if ref.Cast == "" && jsonComparisonOps[op] && jsonValueIsNumeric(value) {
|
||||
ref.Cast = "numeric"
|
||||
}
|
||||
|
||||
colExpr, colArgs := ref.SQL(tableAlias)
|
||||
|
||||
// prepend copies the column-expression args (the bound JSON path, and any
|
||||
// others) ahead of the value args so placeholder order matches the SQL.
|
||||
prepend := func(valueArgs ...interface{}) []interface{} {
|
||||
out := make([]interface{}, 0, len(colArgs)+len(valueArgs))
|
||||
out = append(out, colArgs...)
|
||||
out = append(out, valueArgs...)
|
||||
return out
|
||||
}
|
||||
|
||||
switch op {
|
||||
case "eq", "equals", "=":
|
||||
return fmt.Sprintf("%s = ?", colExpr), prepend(value), true
|
||||
case "neq", "not_equals", "ne", "!=", "<>":
|
||||
return fmt.Sprintf("%s != ?", colExpr), prepend(value), true
|
||||
case "gt", "greater_than", ">":
|
||||
return fmt.Sprintf("%s > ?", colExpr), prepend(value), true
|
||||
case "gte", "greater_than_equals", "ge", ">=":
|
||||
return fmt.Sprintf("%s >= ?", colExpr), prepend(value), true
|
||||
case "lt", "less_than", "<":
|
||||
return fmt.Sprintf("%s < ?", colExpr), prepend(value), true
|
||||
case "lte", "less_than_equals", "le", "<=":
|
||||
return fmt.Sprintf("%s <= ?", colExpr), prepend(value), true
|
||||
case "like":
|
||||
return fmt.Sprintf("%s LIKE ?", colExpr), prepend(value), true
|
||||
case "ilike":
|
||||
return fmt.Sprintf("%s ILIKE ?", colExpr), prepend(value), true
|
||||
case "in":
|
||||
inCond, inArgs := BuildInCondition(colExpr, value)
|
||||
if inCond == "" {
|
||||
return "", nil, false
|
||||
}
|
||||
return inCond, prepend(inArgs...), true
|
||||
case "between", "between_inclusive":
|
||||
lo, hi, bok := twoBoundValues(value)
|
||||
if !bok {
|
||||
return "", nil, false
|
||||
}
|
||||
loOp, hiOp := ">", "<"
|
||||
if op == "between_inclusive" {
|
||||
loOp, hiOp = ">=", "<="
|
||||
}
|
||||
// colExpr appears twice, so its bound args (the JSON path) appear twice.
|
||||
betweenArgs := make([]interface{}, 0, 2*len(colArgs)+2)
|
||||
betweenArgs = append(betweenArgs, colArgs...)
|
||||
betweenArgs = append(betweenArgs, lo)
|
||||
betweenArgs = append(betweenArgs, colArgs...)
|
||||
betweenArgs = append(betweenArgs, hi)
|
||||
return fmt.Sprintf("(%s %s ? AND %s %s ?)", colExpr, loOp, colExpr, hiOp), betweenArgs, true
|
||||
case "is_null", "isnull":
|
||||
return fmt.Sprintf("%s IS NULL", colExpr), prepend(), true
|
||||
case "is_not_null", "isnotnull":
|
||||
return fmt.Sprintf("%s IS NOT NULL", colExpr), prepend(), true
|
||||
default:
|
||||
return "", nil, false
|
||||
}
|
||||
}
|
||||
|
||||
// jsonValueIsNumeric reports whether value (or every element of a 2-slice) is a
|
||||
// number or a numeric-looking string.
|
||||
func jsonValueIsNumeric(value interface{}) bool {
|
||||
switch v := value.(type) {
|
||||
case []interface{}:
|
||||
if len(v) == 0 {
|
||||
return false
|
||||
}
|
||||
for _, e := range v {
|
||||
if !jsonValueIsNumeric(e) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
case []string:
|
||||
if len(v) == 0 {
|
||||
return false
|
||||
}
|
||||
for _, e := range v {
|
||||
if _, ok := toFloat(e); !ok {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
case string:
|
||||
_, ok := toFloat(v)
|
||||
return ok
|
||||
default:
|
||||
_, ok := toFloat(value)
|
||||
return ok
|
||||
}
|
||||
}
|
||||
|
||||
// twoBoundValues extracts the low/high bounds from a BETWEEN filter value.
|
||||
func twoBoundValues(value interface{}) (lo, hi interface{}, ok bool) {
|
||||
switch v := value.(type) {
|
||||
case []interface{}:
|
||||
if len(v) == 2 {
|
||||
return v[0], v[1], true
|
||||
}
|
||||
case []string:
|
||||
if len(v) == 2 {
|
||||
return v[0], v[1], true
|
||||
}
|
||||
}
|
||||
return nil, nil, false
|
||||
}
|
||||
@@ -0,0 +1,164 @@
|
||||
package common
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||
)
|
||||
|
||||
type jsonCondModel struct {
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Data spectypes.SqlJSONB `json:"data"`
|
||||
}
|
||||
|
||||
func TestResolveJSONColumnRef_Gate(t *testing.T) {
|
||||
m := jsonCondModel{}
|
||||
|
||||
// Explicit operator syntax needs no model confirmation.
|
||||
if _, ok := ResolveJSONColumnRef(m, "data->>'city'"); !ok {
|
||||
t.Error("arrow syntax should resolve")
|
||||
}
|
||||
// Dotted shorthand on a real JSON column resolves.
|
||||
if ref, ok := ResolveJSONColumnRef(m, "data.city"); !ok || !reflect.DeepEqual(ref.Path, []string{"city"}) {
|
||||
t.Errorf("dotted shorthand on JSON column should resolve, got ok=%v ref=%+v", ok, ref)
|
||||
}
|
||||
// Dotted shorthand on a non-JSON column must NOT be treated as JSON.
|
||||
if _, ok := ResolveJSONColumnRef(m, "name.first"); ok {
|
||||
t.Error("dotted shorthand on non-JSON column must not resolve as JSON")
|
||||
}
|
||||
// Plain columns never resolve.
|
||||
if _, ok := ResolveJSONColumnRef(m, "name"); ok {
|
||||
t.Error("plain column must not resolve")
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveJSONColumnExpr(t *testing.T) {
|
||||
m := jsonCondModel{}
|
||||
|
||||
expr, args, alias, ok := ResolveJSONColumnExpr(m, "t", "data->'addr'->>'city'")
|
||||
if !ok {
|
||||
t.Fatal("expected ok")
|
||||
}
|
||||
if expr != `("t"."data" #>> ?::text[])` {
|
||||
t.Errorf("expr = %q", expr)
|
||||
}
|
||||
if !reflect.DeepEqual(args, []interface{}{"{addr,city}"}) {
|
||||
t.Errorf("args = %#v", args)
|
||||
}
|
||||
if alias != "data_addr_city" {
|
||||
t.Errorf("alias = %q", alias)
|
||||
}
|
||||
|
||||
if _, _, _, ok := ResolveJSONColumnExpr(m, "t", "name"); ok {
|
||||
t.Error("plain column must not resolve")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildJSONFilterCondition(t *testing.T) {
|
||||
m := jsonCondModel{}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
token string
|
||||
operator string
|
||||
value interface{}
|
||||
wantCond string
|
||||
wantArgs []interface{}
|
||||
}{
|
||||
{
|
||||
name: "eq stays text", token: "data->>'city'", operator: "eq", value: "LA",
|
||||
wantCond: `("data" #>> ?::text[]) = ?`,
|
||||
wantArgs: []interface{}{"{city}", "LA"},
|
||||
},
|
||||
{
|
||||
name: "gt numeric value infers numeric cast", token: "data->>'age'", operator: "gt", value: 18,
|
||||
wantCond: `(("data" #>> ?::text[]))::numeric > ?`,
|
||||
wantArgs: []interface{}{"{age}", 18},
|
||||
},
|
||||
{
|
||||
name: "gt non-numeric value stays text", token: "data->>'name'", operator: "gt", value: "m",
|
||||
wantCond: `("data" #>> ?::text[]) > ?`,
|
||||
wantArgs: []interface{}{"{name}", "m"},
|
||||
},
|
||||
{
|
||||
name: "explicit cast is respected for lt", token: "data->>'ts'::timestamptz", operator: "lt", value: "2020-01-01",
|
||||
wantCond: `(("data" #>> ?::text[]))::timestamptz < ?`,
|
||||
wantArgs: []interface{}{"{ts}", "2020-01-01"},
|
||||
},
|
||||
{
|
||||
name: "ilike", token: "data->>'city'", operator: "ilike", value: "%la%",
|
||||
wantCond: `("data" #>> ?::text[]) ILIKE ?`,
|
||||
wantArgs: []interface{}{"{city}", "%la%"},
|
||||
},
|
||||
{
|
||||
name: "in", token: "data->>'tier'", operator: "in", value: []string{"a", "b"},
|
||||
wantCond: `("data" #>> ?::text[]) IN (?,?)`,
|
||||
wantArgs: []interface{}{"{tier}", "a", "b"},
|
||||
},
|
||||
{
|
||||
name: "between numeric", token: "data->>'age'", operator: "between", value: []interface{}{10, 20},
|
||||
wantCond: `((("data" #>> ?::text[]))::numeric > ? AND (("data" #>> ?::text[]))::numeric < ?)`,
|
||||
wantArgs: []interface{}{"{age}", 10, "{age}", 20},
|
||||
},
|
||||
{
|
||||
name: "is_null", token: "data->>'city'", operator: "is_null", value: nil,
|
||||
wantCond: `("data" #>> ?::text[]) IS NULL`,
|
||||
wantArgs: []interface{}{"{city}"},
|
||||
},
|
||||
{
|
||||
name: "hash path", token: "data#>>'{a,b}'", operator: "eq", value: "x",
|
||||
wantCond: `("data" #>> ?::text[]) = ?`,
|
||||
wantArgs: []interface{}{"{a,b}", "x"},
|
||||
},
|
||||
{
|
||||
name: "dotted shorthand on json column", token: "data.city", operator: "eq", value: "x",
|
||||
wantCond: `("data" #>> ?::text[]) = ?`,
|
||||
wantArgs: []interface{}{"{city}", "x"},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
cond, args, ok := BuildJSONFilterCondition(m, "", tc.token, tc.operator, tc.value)
|
||||
if !ok {
|
||||
t.Fatalf("ok=false for %q", tc.token)
|
||||
}
|
||||
if cond != tc.wantCond {
|
||||
t.Errorf("cond = %q, want %q", cond, tc.wantCond)
|
||||
}
|
||||
if !reflect.DeepEqual(args, tc.wantArgs) {
|
||||
t.Errorf("args = %#v, want %#v", args, tc.wantArgs)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildJSONFilterCondition_NotJSON(t *testing.T) {
|
||||
m := jsonCondModel{}
|
||||
for _, tok := range []string{"name", "id", "name.first"} {
|
||||
if _, _, ok := BuildJSONFilterCondition(m, "", tok, "eq", "x"); ok {
|
||||
t.Errorf("BuildJSONFilterCondition(%q) ok=true, want false", tok)
|
||||
}
|
||||
}
|
||||
// Unknown operator on a real JSON ref -> caller keeps its own handling.
|
||||
if _, _, ok := BuildJSONFilterCondition(m, "", "data->>'x'", "st_intersects", "y"); ok {
|
||||
t.Error("unknown operator must yield ok=false")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildJSONFilterCondition_QualifiedAndInjectionSafe(t *testing.T) {
|
||||
m := jsonCondModel{}
|
||||
// A hostile key never reaches the SQL string — it is bound in the text[] arg.
|
||||
cond, args, ok := BuildJSONFilterCondition(m, "pub.tbl", "data->>'ev\"il'", "eq", "x")
|
||||
if !ok {
|
||||
t.Fatal("ok=false")
|
||||
}
|
||||
if cond != `("pub"."tbl"."data" #>> ?::text[]) = ?` {
|
||||
t.Errorf("cond = %q", cond)
|
||||
}
|
||||
if !reflect.DeepEqual(args, []interface{}{`{"ev\"il"}`, "x"}) {
|
||||
t.Errorf("args = %#v", args)
|
||||
}
|
||||
}
|
||||
@@ -109,6 +109,19 @@ func (v *ColumnValidator) ValidateColumn(column string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// JSON-traversing references (data->>'x', data#>>'{a,b}', or the dotted
|
||||
// data.x shorthand): validate the base column, and for the ambiguous
|
||||
// dotted form require that the base is actually a JSON column.
|
||||
if ref, isJSON := ParseColumnRef(column); isJSON {
|
||||
if ref.Ambiguous && !reflection.IsJSONColumn(v.model, ref.Base) {
|
||||
return fmt.Errorf("invalid column '%s': '%s' is not a JSON column", column, ref.Base)
|
||||
}
|
||||
if _, exists := v.validColumns[strings.ToLower(ref.Base)]; !exists {
|
||||
return fmt.Errorf("invalid column '%s': column does not exist in model", column)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Extract source column name (remove JSON operators like ->> or ->)
|
||||
sourceColumn := reflection.ExtractSourceColumn(column)
|
||||
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"testing"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||
)
|
||||
|
||||
func TestExtractSourceColumn(t *testing.T) {
|
||||
@@ -124,3 +125,35 @@ func TestValidateColumnWithJSONOperators(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateColumn_JSONPathsAndDottedShorthand(t *testing.T) {
|
||||
type Model struct {
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Data spectypes.SqlJSONB `json:"data"`
|
||||
}
|
||||
v := NewColumnValidator(Model{})
|
||||
|
||||
valid := []string{
|
||||
"data->>'city'",
|
||||
"data->'addr'->>'city'",
|
||||
"data#>>'{addr,city}'",
|
||||
"data.addr.city", // dotted shorthand, base is JSON -> allowed
|
||||
"(data->>'age')::int", // cast + paren
|
||||
}
|
||||
for _, c := range valid {
|
||||
if err := v.ValidateColumn(c); err != nil {
|
||||
t.Errorf("ValidateColumn(%q) = %v, want nil", c, err)
|
||||
}
|
||||
}
|
||||
|
||||
invalid := []string{
|
||||
"nope->>'city'", // base column does not exist
|
||||
"name.first", // dotted shorthand but 'name' is not a JSON column
|
||||
}
|
||||
for _, c := range invalid {
|
||||
if err := v.ValidateColumn(c); err == nil {
|
||||
t.Errorf("ValidateColumn(%q) = nil, want error", c)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+28
-4
@@ -676,7 +676,7 @@ func (h *Handler) readByID(hookCtx *HookContext) (interface{}, error) {
|
||||
|
||||
// Apply columns
|
||||
if hookCtx.Options != nil && len(hookCtx.Options.Columns) > 0 {
|
||||
query = query.Column(hookCtx.Options.Columns...)
|
||||
query = common.ApplySelectColumns(query, hookCtx.Model, "", hookCtx.Options.Columns)
|
||||
}
|
||||
|
||||
// Apply preloads (simplified)
|
||||
@@ -714,9 +714,19 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata
|
||||
if hookCtx.Options != nil {
|
||||
// Apply filters
|
||||
for _, filter := range hookCtx.Options.Filters {
|
||||
if cond, jargs, ok := common.BuildJSONFilterCondition(hookCtx.Model, "", filter.Column, filter.Operator, filter.Value); ok {
|
||||
query = query.Where(cond, jargs...)
|
||||
continue
|
||||
}
|
||||
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)
|
||||
}
|
||||
@@ -728,6 +738,10 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata
|
||||
if sort.Direction == "desc" {
|
||||
direction = "DESC"
|
||||
}
|
||||
if expr, jargs, _, ok := common.ResolveJSONColumnExpr(hookCtx.Model, "", sort.Column); ok {
|
||||
query = query.OrderExpr(fmt.Sprintf("%s %s", expr, direction), jargs...)
|
||||
continue
|
||||
}
|
||||
query = query.Order(fmt.Sprintf("%s %s", sort.Column, direction))
|
||||
}
|
||||
|
||||
@@ -746,7 +760,7 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata
|
||||
|
||||
// Apply columns
|
||||
if len(hookCtx.Options.Columns) > 0 {
|
||||
query = query.Column(hookCtx.Options.Columns...)
|
||||
query = common.ApplySelectColumns(query, hookCtx.Model, "", hookCtx.Options.Columns)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -772,9 +786,19 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata
|
||||
countQuery := h.db.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName)
|
||||
if hookCtx.Options != nil {
|
||||
for _, filter := range hookCtx.Options.Filters {
|
||||
if cond, jargs, ok := common.BuildJSONFilterCondition(hookCtx.Model, "", filter.Column, filter.Operator, filter.Value); ok {
|
||||
countQuery = countQuery.Where(cond, jargs...)
|
||||
continue
|
||||
}
|
||||
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)
|
||||
}
|
||||
|
||||
@@ -9,6 +9,7 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||
)
|
||||
|
||||
type PrimaryKeyNameProvider interface {
|
||||
@@ -728,19 +729,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)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -3,6 +3,8 @@ package reflection
|
||||
import (
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||
)
|
||||
|
||||
// Test models for GORM
|
||||
@@ -1047,6 +1049,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
|
||||
|
||||
@@ -1,17 +1,19 @@
|
||||
package reflection
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"reflect"
|
||||
"strings"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||
)
|
||||
|
||||
// getColumnFieldType resolves the reflect.Type of the struct field that backs
|
||||
// colName (matched by json tag, field name or snake_case), following the same
|
||||
// rules as GetColumnTypeFromModel.
|
||||
func getColumnFieldType(model interface{}, colName string) (reflect.Type, bool) {
|
||||
// getColumnStructField resolves the struct field that backs colName (matched by
|
||||
// json tag, field name or snake_case), following the same rules as
|
||||
// GetColumnTypeFromModel.
|
||||
func getColumnStructField(model interface{}, colName string) (reflect.StructField, bool) {
|
||||
if model == nil {
|
||||
return nil, false
|
||||
return reflect.StructField{}, false
|
||||
}
|
||||
sourceColName := ExtractSourceColumn(colName)
|
||||
|
||||
@@ -20,7 +22,7 @@ func getColumnFieldType(model interface{}, colName string) (reflect.Type, bool)
|
||||
modelType = modelType.Elem()
|
||||
}
|
||||
if modelType == nil || modelType.Kind() != reflect.Struct {
|
||||
return nil, false
|
||||
return reflect.StructField{}, false
|
||||
}
|
||||
|
||||
for i := 0; i < modelType.NumField(); i++ {
|
||||
@@ -28,17 +30,28 @@ func getColumnFieldType(model interface{}, colName string) (reflect.Type, bool)
|
||||
|
||||
if jsonTag := field.Tag.Get("json"); jsonTag != "" {
|
||||
if name := jsonTagName(jsonTag); name == sourceColName {
|
||||
return field.Type, true
|
||||
return field, true
|
||||
}
|
||||
}
|
||||
if equalFold(field.Name, sourceColName) {
|
||||
return field.Type, true
|
||||
return field, true
|
||||
}
|
||||
if ToSnakeCase(field.Name) == sourceColName {
|
||||
return field.Type, true
|
||||
return field, true
|
||||
}
|
||||
}
|
||||
return nil, false
|
||||
return reflect.StructField{}, false
|
||||
}
|
||||
|
||||
// getColumnFieldType resolves the reflect.Type of the struct field that backs
|
||||
// colName (matched by json tag, field name or snake_case), following the same
|
||||
// rules as GetColumnTypeFromModel.
|
||||
func getColumnFieldType(model interface{}, colName string) (reflect.Type, bool) {
|
||||
f, ok := getColumnStructField(model, colName)
|
||||
if !ok {
|
||||
return nil, false
|
||||
}
|
||||
return f.Type, true
|
||||
}
|
||||
|
||||
func jsonTagName(tag string) string {
|
||||
@@ -92,3 +105,80 @@ func IsVectorColumn(model interface{}, colName string) bool {
|
||||
t, ok := getColumnFieldType(model, colName)
|
||||
return ok && spectypes.IsVectorType(t)
|
||||
}
|
||||
|
||||
var rawMessageType = reflect.TypeOf(json.RawMessage(nil))
|
||||
|
||||
// IsJSONColumn reports whether colName is backed by a JSON/JSONB column on the
|
||||
// model. It recognises the spectypes SqlJSONB wrapper, encoding/json.RawMessage,
|
||||
// map-typed fields, and fields carrying a bun/gorm `type:json` / `type:jsonb`
|
||||
// tag. colName should be a bare column name (callers pass the parsed base column
|
||||
// of a JSON path, not the full "col->>'x'" expression).
|
||||
func IsJSONColumn(model interface{}, colName string) bool {
|
||||
f, ok := getColumnStructField(model, colName)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
|
||||
ft := f.Type
|
||||
for ft != nil && ft.Kind() == reflect.Pointer {
|
||||
ft = ft.Elem()
|
||||
}
|
||||
if ft == nil {
|
||||
return false
|
||||
}
|
||||
|
||||
if spectypes.IsJSONType(ft) {
|
||||
return true
|
||||
}
|
||||
if ft == rawMessageType {
|
||||
return true
|
||||
}
|
||||
if ft.Kind() == reflect.Map {
|
||||
return true
|
||||
}
|
||||
if tagDeclaresJSON(f.Tag.Get("bun")) || tagDeclaresJSON(f.Tag.Get("gorm")) {
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// 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 ""
|
||||
}
|
||||
for _, part := range strings.FieldsFunc(tag, func(r rune) bool {
|
||||
return r == ',' || r == ';' || r == ' '
|
||||
}) {
|
||||
value, found := strings.CutPrefix(strings.TrimSpace(part), "type:")
|
||||
if !found {
|
||||
continue
|
||||
}
|
||||
return strings.ToLower(strings.TrimSpace(value))
|
||||
}
|
||||
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"
|
||||
}
|
||||
|
||||
@@ -0,0 +1,59 @@
|
||||
package reflection
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||
)
|
||||
|
||||
type jsonColModel struct {
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Meta spectypes.SqlJSONB `json:"meta"`
|
||||
Raw json.RawMessage `json:"raw"`
|
||||
Attrs map[string]interface{} `json:"attrs"`
|
||||
Config []byte `json:"config" bun:"config,type:jsonb"`
|
||||
Settings string `json:"settings" gorm:"column:settings;type:json"`
|
||||
Blob []byte `json:"blob"`
|
||||
}
|
||||
|
||||
func TestIsJSONColumn(t *testing.T) {
|
||||
m := jsonColModel{}
|
||||
|
||||
jsonCols := []string{"meta", "raw", "attrs", "config", "settings"}
|
||||
for _, c := range jsonCols {
|
||||
if !IsJSONColumn(m, c) {
|
||||
t.Errorf("expected %q to be a JSON column", c)
|
||||
}
|
||||
}
|
||||
|
||||
notJSON := []string{"id", "name", "blob", "missing"}
|
||||
for _, c := range notJSON {
|
||||
if IsJSONColumn(m, c) {
|
||||
t.Errorf("expected %q NOT to be a JSON column", c)
|
||||
}
|
||||
}
|
||||
|
||||
if IsJSONColumn(nil, "meta") {
|
||||
t.Error("nil model must not report JSON columns")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTagDeclaresJSON(t *testing.T) {
|
||||
cases := map[string]bool{
|
||||
"config,type:jsonb": true,
|
||||
"column:settings;type:json": true,
|
||||
"col,type:text": false,
|
||||
"column:name": false,
|
||||
"": false,
|
||||
"col,type:jsonb,notnull": true,
|
||||
"column:x;type:varchar(255)": false,
|
||||
"col , type:json": true,
|
||||
}
|
||||
for tag, want := range cases {
|
||||
if got := tagDeclaresJSON(tag); got != want {
|
||||
t.Errorf("tagDeclaresJSON(%q) = %v; want %v", tag, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
+19
-10
@@ -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
|
||||
|
||||
@@ -128,7 +128,7 @@ func TestBuildFilterCondition(t *testing.T) {
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
condition, args := h.buildFilterCondition(tt.filter)
|
||||
condition, args := h.buildFilterCondition(tt.filter, nil)
|
||||
|
||||
if condition != tt.expectedCondition {
|
||||
t.Errorf("Expected condition '%s', got '%s'", tt.expectedCondition, condition)
|
||||
|
||||
+162
-117
@@ -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"
|
||||
@@ -347,6 +348,10 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
||||
if len(options.Columns) > 0 {
|
||||
logger.Debug("Selecting columns: %v", options.Columns)
|
||||
for _, col := range options.Columns {
|
||||
if expr, jargs, alias, ok := common.ResolveJSONColumnExpr(model, "", col); ok {
|
||||
query = query.ColumnExpr(expr+" AS "+common.QuoteIdent(alias), jargs...)
|
||||
continue
|
||||
}
|
||||
query = query.Column(reflection.ExtractSourceColumn(col))
|
||||
}
|
||||
}
|
||||
@@ -393,7 +398,7 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
||||
}
|
||||
|
||||
// Apply filters with proper grouping for OR logic
|
||||
query = h.applyFilters(query, options.Filters)
|
||||
query = h.applyFilters(query, options.Filters, model)
|
||||
|
||||
// Apply custom operators
|
||||
for _, customOp := range options.CustomOperators {
|
||||
@@ -413,6 +418,10 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
||||
direction = "DESC"
|
||||
}
|
||||
logger.Debug("Applying sort: %s %s", sort.Column, direction)
|
||||
if expr, jargs, _, ok := common.ResolveJSONColumnExpr(model, "", sort.Column); ok {
|
||||
query = query.OrderExpr(fmt.Sprintf("%s %s", expr, direction), jargs...)
|
||||
continue
|
||||
}
|
||||
query = query.Order(fmt.Sprintf("%s %s", sort.Column, direction))
|
||||
}
|
||||
|
||||
@@ -536,7 +545,7 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
||||
|
||||
// Apply the same filters as the main query
|
||||
for _, filter := range options.Filters {
|
||||
rowNumQuery = h.applyFilter(rowNumQuery, filter)
|
||||
rowNumQuery = h.applyFilter(rowNumQuery, filter, model)
|
||||
}
|
||||
|
||||
// Apply custom operators
|
||||
@@ -714,15 +723,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)
|
||||
@@ -761,15 +771,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)
|
||||
@@ -843,15 +854,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)
|
||||
@@ -890,15 +902,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)
|
||||
@@ -974,15 +987,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)
|
||||
@@ -1027,15 +1041,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)
|
||||
@@ -1158,16 +1173,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 {
|
||||
@@ -1379,16 +1395,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 {
|
||||
@@ -1535,16 +1552,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 {
|
||||
@@ -1640,15 +1658,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)
|
||||
@@ -1829,7 +1848,7 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id
|
||||
// applyFilters applies all filters with proper grouping for OR logic
|
||||
// Groups consecutive OR filters together to ensure proper query precedence
|
||||
// Example: [A, B(OR), C(OR), D(AND)] => WHERE (A OR B OR C) AND D
|
||||
func (h *Handler) applyFilters(query common.SelectQuery, filters []common.FilterOption) common.SelectQuery {
|
||||
func (h *Handler) applyFilters(query common.SelectQuery, filters []common.FilterOption, model interface{}) common.SelectQuery {
|
||||
if len(filters) == 0 {
|
||||
return query
|
||||
}
|
||||
@@ -1849,11 +1868,11 @@ func (h *Handler) applyFilters(query common.SelectQuery, filters []common.Filter
|
||||
}
|
||||
|
||||
// Apply the OR group as a single grouped WHERE clause
|
||||
query = h.applyFilterGroup(query, orGroup)
|
||||
query = h.applyFilterGroup(query, orGroup, model)
|
||||
i = j
|
||||
} else {
|
||||
// Single filter with AND logic (or first filter)
|
||||
condition, args := h.buildFilterCondition(filters[i])
|
||||
condition, args := h.buildFilterCondition(filters[i], model)
|
||||
if condition != "" {
|
||||
query = query.Where(condition, args...)
|
||||
}
|
||||
@@ -1866,7 +1885,7 @@ func (h *Handler) applyFilters(query common.SelectQuery, filters []common.Filter
|
||||
|
||||
// applyFilterGroup applies a group of filters that should be OR'd together
|
||||
// Always wraps them in parentheses and applies as a single WHERE clause
|
||||
func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.FilterOption) common.SelectQuery {
|
||||
func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.FilterOption, model interface{}) common.SelectQuery {
|
||||
if len(filters) == 0 {
|
||||
return query
|
||||
}
|
||||
@@ -1876,7 +1895,7 @@ func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.Fi
|
||||
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...)
|
||||
@@ -1897,11 +1916,18 @@ func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.Fi
|
||||
return query.Where(groupedCondition, args...)
|
||||
}
|
||||
|
||||
// buildFilterCondition builds a filter condition and returns it with args
|
||||
func (h *Handler) buildFilterCondition(filter common.FilterOption) (conditionString string, conditionArgs []interface{}) {
|
||||
// buildFilterCondition builds a filter condition and returns it with args.
|
||||
// model, when non-nil, lets JSON sub-field references (data->>'x', data#>>'{a,b}',
|
||||
// or the dotted data.x shorthand for a JSON column) resolve to a safe,
|
||||
// parameterised expression before the ordinary operator handling below.
|
||||
func (h *Handler) buildFilterCondition(filter common.FilterOption, model interface{}) (conditionString string, conditionArgs []interface{}) {
|
||||
var condition string
|
||||
var args []interface{}
|
||||
|
||||
if cond, jargs, ok := common.BuildJSONFilterCondition(model, "", filter.Column, filter.Operator, filter.Value); ok {
|
||||
return cond, jargs
|
||||
}
|
||||
|
||||
switch filter.Operator {
|
||||
case "eq", "=":
|
||||
condition = fmt.Sprintf("%s = ?", filter.Column)
|
||||
@@ -1922,10 +1948,10 @@ func (h *Handler) buildFilterCondition(filter common.FilterOption) (conditionStr
|
||||
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)
|
||||
@@ -1958,13 +1984,32 @@ func (h *Handler) buildFilterCondition(filter common.FilterOption) (conditionStr
|
||||
return condition, args
|
||||
}
|
||||
|
||||
func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOption) common.SelectQuery {
|
||||
// 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")
|
||||
|
||||
var condition string
|
||||
var args []interface{}
|
||||
|
||||
if cond, jargs, ok := common.BuildJSONFilterCondition(model, "", filter.Column, filter.Operator, filter.Value); ok {
|
||||
if useOrLogic {
|
||||
return query.WhereOr(cond, jargs...)
|
||||
}
|
||||
return query.Where(cond, jargs...)
|
||||
}
|
||||
|
||||
switch filter.Operator {
|
||||
case "eq", "=":
|
||||
condition = fmt.Sprintf("%s = ?", filter.Column)
|
||||
@@ -1985,10 +2030,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)
|
||||
@@ -2394,7 +2439,7 @@ func (h *Handler) applyPreloads(model interface{}, query common.SelectQuery, pre
|
||||
|
||||
if len(preload.Filters) > 0 {
|
||||
for _, filter := range preload.Filters {
|
||||
sq = h.applyFilter(sq, filter)
|
||||
sq = h.applyFilter(sq, filter, nil)
|
||||
}
|
||||
}
|
||||
if len(preload.Sort) > 0 {
|
||||
|
||||
@@ -0,0 +1,153 @@
|
||||
package resolvespec
|
||||
|
||||
import (
|
||||
"context"
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||
)
|
||||
|
||||
// jsonColModel has a real JSONB column so the dotted "data.x" shorthand is
|
||||
// recognised as JSON access.
|
||||
type jsonColModel struct {
|
||||
ID int64 `json:"id" bun:"id,pk"`
|
||||
Name string `json:"name" bun:"name"`
|
||||
Data spectypes.SqlJSONB `json:"data" bun:"data"`
|
||||
}
|
||||
|
||||
type jsonCapCall struct {
|
||||
method string
|
||||
query string
|
||||
args []interface{}
|
||||
}
|
||||
|
||||
// jsonCapQuery records the string + args of the calls the handler makes.
|
||||
type jsonCapQuery struct {
|
||||
calls []jsonCapCall
|
||||
}
|
||||
|
||||
func (m *jsonCapQuery) rec(method, query string, args []interface{}) common.SelectQuery {
|
||||
m.calls = append(m.calls, jsonCapCall{method: method, query: query, args: args})
|
||||
return m
|
||||
}
|
||||
|
||||
func (m *jsonCapQuery) Model(interface{}) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) Table(string) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) Column(cols ...string) common.SelectQuery {
|
||||
for _, c := range cols {
|
||||
m.rec("Column", c, nil)
|
||||
}
|
||||
return m
|
||||
}
|
||||
func (m *jsonCapQuery) ColumnExpr(q string, args ...interface{}) common.SelectQuery {
|
||||
return m.rec("ColumnExpr", q, args)
|
||||
}
|
||||
func (m *jsonCapQuery) Where(q string, args ...interface{}) common.SelectQuery {
|
||||
return m.rec("Where", q, args)
|
||||
}
|
||||
func (m *jsonCapQuery) WhereOr(q string, args ...interface{}) common.SelectQuery {
|
||||
return m.rec("WhereOr", q, args)
|
||||
}
|
||||
func (m *jsonCapQuery) Join(string, ...interface{}) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) LeftJoin(string, ...interface{}) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) Preload(string, ...interface{}) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) PreloadRelation(string, ...func(common.SelectQuery) common.SelectQuery) common.SelectQuery {
|
||||
return m
|
||||
}
|
||||
func (m *jsonCapQuery) JoinRelation(string, ...func(common.SelectQuery) common.SelectQuery) common.SelectQuery {
|
||||
return m
|
||||
}
|
||||
func (m *jsonCapQuery) Order(o string) common.SelectQuery { return m.rec("Order", o, nil) }
|
||||
func (m *jsonCapQuery) OrderExpr(o string, args ...interface{}) common.SelectQuery {
|
||||
return m.rec("OrderExpr", o, args)
|
||||
}
|
||||
func (m *jsonCapQuery) Limit(int) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) Offset(int) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) Group(string) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) Having(string, ...interface{}) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) Scan(context.Context, interface{}) error { return nil }
|
||||
func (m *jsonCapQuery) ScanModel(context.Context) error { return nil }
|
||||
func (m *jsonCapQuery) Count(context.Context) (int, error) { return 0, nil }
|
||||
func (m *jsonCapQuery) Exists(context.Context) (bool, error) { return false, nil }
|
||||
|
||||
func (m *jsonCapQuery) only(t *testing.T) jsonCapCall {
|
||||
t.Helper()
|
||||
if len(m.calls) != 1 {
|
||||
t.Fatalf("expected exactly 1 recorded call, got %d: %+v", len(m.calls), m.calls)
|
||||
}
|
||||
return m.calls[0]
|
||||
}
|
||||
|
||||
func TestBuildFilterCondition_JSONColumn(t *testing.T) {
|
||||
h := &Handler{}
|
||||
model := jsonColModel{}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
filter common.FilterOption
|
||||
wantCond string
|
||||
wantArgs []interface{}
|
||||
}{
|
||||
{
|
||||
name: "arrow syntax eq stays text",
|
||||
filter: common.FilterOption{Column: "data->>'city'", Operator: "eq", Value: "LA"},
|
||||
wantCond: `("data" #>> ?::text[]) = ?`,
|
||||
wantArgs: []interface{}{"{city}", "LA"},
|
||||
},
|
||||
{
|
||||
name: "dotted shorthand numeric cast inference",
|
||||
filter: common.FilterOption{Column: "data.age", Operator: "gt", Value: 18},
|
||||
wantCond: `(("data" #>> ?::text[]))::numeric > ?`,
|
||||
wantArgs: []interface{}{"{age}", 18},
|
||||
},
|
||||
{
|
||||
name: "hash path with explicit cast",
|
||||
filter: common.FilterOption{Column: "data#>>'{a,b}'::int", Operator: "lte", Value: "5"},
|
||||
wantCond: `(("data" #>> ?::text[]))::integer <= ?`,
|
||||
wantArgs: []interface{}{"{a,b}", "5"},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
cond, args := h.buildFilterCondition(tc.filter, model)
|
||||
if cond != tc.wantCond {
|
||||
t.Fatalf("cond = %q, want %q", cond, tc.wantCond)
|
||||
}
|
||||
if !reflect.DeepEqual(args, tc.wantArgs) {
|
||||
t.Fatalf("args = %#v, want %#v", args, tc.wantArgs)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// Non-JSON column falls through to ordinary handling.
|
||||
cond, _ := h.buildFilterCondition(common.FilterOption{Column: "name", Operator: "eq", Value: "x"}, model)
|
||||
if cond != "name = ?" {
|
||||
t.Fatalf("non-JSON cond = %q", cond)
|
||||
}
|
||||
|
||||
// Without a model the dotted shorthand must NOT be treated as JSON.
|
||||
cond, _ = h.buildFilterCondition(common.FilterOption{Column: "data.age", Operator: "eq", Value: "x"}, nil)
|
||||
if cond != "data.age = ?" {
|
||||
t.Fatalf("nil-model dotted cond = %q, want ordinary handling", cond)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyFilter_JSONColumn(t *testing.T) {
|
||||
h := &Handler{}
|
||||
model := jsonColModel{}
|
||||
|
||||
q := &jsonCapQuery{}
|
||||
h.applyFilter(q, common.FilterOption{
|
||||
Column: "data->>'tier'", Operator: "in", Value: []string{"a", "b"}, LogicOperator: "OR",
|
||||
}, model)
|
||||
c := q.only(t)
|
||||
if c.method != "WhereOr" || c.query != `("data" #>> ?::text[]) IN (?,?)` {
|
||||
t.Fatalf("call = %+v", c)
|
||||
}
|
||||
if !reflect.DeepEqual(c.args, []interface{}{"{tier}", "a", "b"}) {
|
||||
t.Fatalf("args = %#v", c.args)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
+123
-20
@@ -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,
|
||||
@@ -471,7 +525,14 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
||||
// Apply column selection
|
||||
if len(options.Columns) > 0 {
|
||||
logger.Debug("Selecting columns: %v", options.Columns)
|
||||
selectAlias := reflection.ExtractTableNameOnly(tableName)
|
||||
for _, col := range options.Columns {
|
||||
// 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 {
|
||||
query = query.ColumnExpr(expr+" AS "+common.QuoteIdent(alias), jargs...)
|
||||
continue
|
||||
}
|
||||
query = query.Column(reflection.ExtractSourceColumn(col))
|
||||
}
|
||||
|
||||
@@ -610,12 +671,12 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
||||
|
||||
// Apply the OR group as a single grouped condition
|
||||
logger.Debug("Applying OR filter group with %d conditions", len(orFilters))
|
||||
query = h.applyOrFilterGroup(query, orFilters, orCastInfo, tableName)
|
||||
query = h.applyOrFilterGroup(query, orFilters, orCastInfo, tableName, model)
|
||||
i = j
|
||||
} else {
|
||||
// Single AND filter - apply normally
|
||||
logger.Debug("Applying filter: %s %s %v (needsCast=%v, logic=%s)", filter.Column, filter.Operator, filter.Value, castInfo.NeedsCast, logicOp)
|
||||
query = h.applyFilter(query, *filter, tableName, castInfo.NeedsCast, logicOp)
|
||||
query = h.applyFilter(query, *filter, tableName, castInfo.NeedsCast, logicOp, model)
|
||||
i++
|
||||
}
|
||||
}
|
||||
@@ -715,8 +776,12 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
||||
}
|
||||
logger.Debug("Applying sort: %s %s", sort.Column, direction)
|
||||
|
||||
// Check if it's an expression (enclosed in brackets) - use directly without quoting
|
||||
if strings.HasPrefix(sort.Column, "(") && strings.HasSuffix(sort.Column, ")") {
|
||||
// JSON sub-field reference (data->>'x', data#>>'{a,b}', or dotted
|
||||
// shorthand when the base is a JSON column) - resolve to a safe
|
||||
// parameterised expression before the generic branches.
|
||||
if expr, jargs, _, ok := common.ResolveJSONColumnExpr(model, tableAlias, sort.Column); ok {
|
||||
query = query.OrderExpr(fmt.Sprintf("%s %s", expr, direction), jargs...)
|
||||
} else if strings.HasPrefix(sort.Column, "(") && strings.HasSuffix(sort.Column, ")") {
|
||||
// For expressions, pass as raw SQL to prevent auto-quoting
|
||||
query = query.OrderExpr(fmt.Sprintf("%s %s", sort.Column, direction))
|
||||
} else if strings.Contains(sort.Column, ".") {
|
||||
@@ -1063,7 +1128,7 @@ func (h *Handler) applyPreloadWithRecursion(query common.SelectQuery, preload co
|
||||
// Apply filters
|
||||
if len(preload.Filters) > 0 {
|
||||
for _, filter := range preload.Filters {
|
||||
sq = h.applyFilter(sq, filter, "", false, "AND")
|
||||
sq = h.applyFilter(sq, filter, "", false, "AND", nil)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1225,6 +1290,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,
|
||||
@@ -1324,6 +1390,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,
|
||||
@@ -1478,6 +1545,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,
|
||||
@@ -1675,6 +1743,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,
|
||||
@@ -1749,6 +1818,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,
|
||||
@@ -1807,6 +1877,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,
|
||||
@@ -1891,6 +1962,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,
|
||||
@@ -2288,16 +2360,11 @@ func (h *Handler) qualifyColumnName(columnName, fullTableName string) string {
|
||||
return fmt.Sprintf("%s.%s", tableOnly, columnName)
|
||||
}
|
||||
|
||||
func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOption, tableName string, needsCast bool, logicOp string) common.SelectQuery {
|
||||
func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOption, tableName string, needsCast bool, logicOp string, model interface{}) common.SelectQuery {
|
||||
// Qualify the column name with table name if not already qualified
|
||||
rawQualifiedColumn := h.qualifyColumnName(filter.Column, tableName)
|
||||
qualifiedColumn := rawQualifiedColumn
|
||||
|
||||
// Apply casting to text if needed for non-numeric columns or non-numeric values
|
||||
if needsCast {
|
||||
qualifiedColumn = fmt.Sprintf("CAST(%s AS TEXT)", rawQualifiedColumn)
|
||||
}
|
||||
|
||||
// Helper function to apply the correct Where method based on logic operator
|
||||
applyWhere := func(condition string, args ...interface{}) common.SelectQuery {
|
||||
if logicOp == "OR" {
|
||||
@@ -2306,6 +2373,26 @@ func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOpti
|
||||
return query.Where(condition, args...)
|
||||
}
|
||||
|
||||
// JSON sub-field access (data->>'x', data#>>'{a,b}', or the dotted data.x
|
||||
// shorthand when "data" is a JSON column): resolve to a safe, parameterised
|
||||
// expression before the ordinary column handling below.
|
||||
tableAlias := reflection.ExtractTableNameOnly(tableName)
|
||||
if cond, jargs, ok := common.BuildJSONFilterCondition(model, tableAlias, filter.Column, filter.Operator, filter.Value); ok {
|
||||
return applyWhere(cond, jargs...)
|
||||
}
|
||||
|
||||
// Apply casting to text if needed for non-numeric columns or non-numeric values
|
||||
if needsCast {
|
||||
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)
|
||||
@@ -2320,11 +2407,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 == "" {
|
||||
@@ -2377,24 +2467,37 @@ func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOpti
|
||||
|
||||
// applyOrFilterGroup applies a group of OR filters as a single grouped condition
|
||||
// This ensures OR conditions are properly grouped with parentheses to prevent OR logic from escaping
|
||||
func (h *Handler) applyOrFilterGroup(query common.SelectQuery, filters []*common.FilterOption, castInfo []ColumnCastInfo, tableName string) common.SelectQuery {
|
||||
func (h *Handler) applyOrFilterGroup(query common.SelectQuery, filters []*common.FilterOption, castInfo []ColumnCastInfo, tableName string, model interface{}) common.SelectQuery {
|
||||
if len(filters) == 0 {
|
||||
return query
|
||||
}
|
||||
|
||||
tableAlias := reflection.ExtractTableNameOnly(tableName)
|
||||
|
||||
// Build individual filter conditions
|
||||
conditions := []string{}
|
||||
args := []interface{}{}
|
||||
|
||||
for i, filter := range filters {
|
||||
// JSON sub-field access: resolve to a safe parameterised condition first.
|
||||
if cond, jargs, ok := common.BuildJSONFilterCondition(model, tableAlias, filter.Column, filter.Operator, filter.Value); ok {
|
||||
conditions = append(conditions, cond)
|
||||
args = append(args, jargs...)
|
||||
continue
|
||||
}
|
||||
|
||||
// Qualify the column name with table name if not already qualified
|
||||
rawQualifiedColumn := h.qualifyColumnName(filter.Column, tableName)
|
||||
qualifiedColumn := rawQualifiedColumn
|
||||
|
||||
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)
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -0,0 +1,149 @@
|
||||
package restheadspec
|
||||
|
||||
import (
|
||||
"context"
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||
)
|
||||
|
||||
// jsonColModel exercises the JSON-column wiring: Data is a real JSONB column so
|
||||
// the dotted "data.x" shorthand is recognised as JSON access.
|
||||
type jsonColModel struct {
|
||||
ID int64 `json:"id" bun:"id,pk"`
|
||||
Name string `json:"name" bun:"name"`
|
||||
Data spectypes.SqlJSONB `json:"data" bun:"data"`
|
||||
}
|
||||
|
||||
// jsonCapQuery is a minimal common.SelectQuery that records the string + args of
|
||||
// the calls the handler makes so a test can assert on them.
|
||||
type jsonCapQuery struct {
|
||||
calls []jsonCapCall
|
||||
}
|
||||
|
||||
type jsonCapCall struct {
|
||||
method string
|
||||
query string
|
||||
args []interface{}
|
||||
}
|
||||
|
||||
func (m *jsonCapQuery) rec(method, query string, args []interface{}) common.SelectQuery {
|
||||
m.calls = append(m.calls, jsonCapCall{method: method, query: query, args: args})
|
||||
return m
|
||||
}
|
||||
|
||||
func (m *jsonCapQuery) Model(interface{}) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) Table(string) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) Column(cols ...string) common.SelectQuery {
|
||||
for _, c := range cols {
|
||||
m.rec("Column", c, nil)
|
||||
}
|
||||
return m
|
||||
}
|
||||
func (m *jsonCapQuery) ColumnExpr(q string, args ...interface{}) common.SelectQuery {
|
||||
return m.rec("ColumnExpr", q, args)
|
||||
}
|
||||
func (m *jsonCapQuery) Where(q string, args ...interface{}) common.SelectQuery {
|
||||
return m.rec("Where", q, args)
|
||||
}
|
||||
func (m *jsonCapQuery) WhereOr(q string, args ...interface{}) common.SelectQuery {
|
||||
return m.rec("WhereOr", q, args)
|
||||
}
|
||||
func (m *jsonCapQuery) WhereIn(col string, values interface{}) common.SelectQuery {
|
||||
return m.rec("WhereIn", col, []interface{}{values})
|
||||
}
|
||||
func (m *jsonCapQuery) Order(o string) common.SelectQuery { return m.rec("Order", o, nil) }
|
||||
func (m *jsonCapQuery) OrderExpr(o string, args ...interface{}) common.SelectQuery {
|
||||
return m.rec("OrderExpr", o, args)
|
||||
}
|
||||
func (m *jsonCapQuery) Limit(int) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) Offset(int) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) Join(string, ...interface{}) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) LeftJoin(string, ...interface{}) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) Group(string) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) Having(string, ...interface{}) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) Preload(string, ...interface{}) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) PreloadRelation(string, ...func(common.SelectQuery) common.SelectQuery) common.SelectQuery {
|
||||
return m
|
||||
}
|
||||
func (m *jsonCapQuery) JoinRelation(string, ...func(common.SelectQuery) common.SelectQuery) common.SelectQuery {
|
||||
return m
|
||||
}
|
||||
func (m *jsonCapQuery) Scan(context.Context, interface{}) error { return nil }
|
||||
func (m *jsonCapQuery) ScanModel(context.Context) error { return nil }
|
||||
func (m *jsonCapQuery) Count(context.Context) (int, error) { return 0, nil }
|
||||
func (m *jsonCapQuery) Exists(context.Context) (bool, error) { return false, nil }
|
||||
func (m *jsonCapQuery) GetUnderlyingQuery() interface{} { return nil }
|
||||
func (m *jsonCapQuery) GetModel() interface{} { return nil }
|
||||
|
||||
func (m *jsonCapQuery) only(t *testing.T) jsonCapCall {
|
||||
t.Helper()
|
||||
if len(m.calls) != 1 {
|
||||
t.Fatalf("expected exactly 1 recorded call, got %d: %+v", len(m.calls), m.calls)
|
||||
}
|
||||
return m.calls[0]
|
||||
}
|
||||
|
||||
func TestApplyFilter_JSONColumn(t *testing.T) {
|
||||
h := &Handler{}
|
||||
model := jsonColModel{}
|
||||
|
||||
t.Run("arrow syntax eq", func(t *testing.T) {
|
||||
q := &jsonCapQuery{}
|
||||
h.applyFilter(q, common.FilterOption{
|
||||
Column: "data->>'city'", Operator: "eq", Value: "LA",
|
||||
}, "public.things", false, "AND", model)
|
||||
c := q.only(t)
|
||||
if c.method != "Where" || c.query != `("things"."data" #>> ?::text[]) = ?` {
|
||||
t.Fatalf("call = %+v", c)
|
||||
}
|
||||
if !reflect.DeepEqual(c.args, []interface{}{"{city}", "LA"}) {
|
||||
t.Fatalf("args = %#v", c.args)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("dotted shorthand with numeric cast inference, OR logic", func(t *testing.T) {
|
||||
q := &jsonCapQuery{}
|
||||
h.applyFilter(q, common.FilterOption{
|
||||
Column: "data.age", Operator: "gt", Value: 18,
|
||||
}, "public.things", false, "OR", model)
|
||||
c := q.only(t)
|
||||
if c.method != "WhereOr" || c.query != `(("things"."data" #>> ?::text[]))::numeric > ?` {
|
||||
t.Fatalf("call = %+v", c)
|
||||
}
|
||||
if !reflect.DeepEqual(c.args, []interface{}{"{age}", 18}) {
|
||||
t.Fatalf("args = %#v", c.args)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("non-JSON column is untouched", func(t *testing.T) {
|
||||
q := &jsonCapQuery{}
|
||||
h.applyFilter(q, common.FilterOption{
|
||||
Column: "name", Operator: "eq", Value: "x",
|
||||
}, "public.things", false, "AND", model)
|
||||
c := q.only(t)
|
||||
if c.query != "things.name = ?" {
|
||||
t.Fatalf("call = %+v", c)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("nil model: explicit syntax still works, dotted does not", func(t *testing.T) {
|
||||
q := &jsonCapQuery{}
|
||||
h.applyFilter(q, common.FilterOption{
|
||||
Column: "data->>'city'", Operator: "eq", Value: "LA",
|
||||
}, "public.things", false, "AND", nil)
|
||||
if c := q.only(t); c.query != `("things"."data" #>> ?::text[]) = ?` {
|
||||
t.Fatalf("explicit call = %+v", c)
|
||||
}
|
||||
|
||||
q2 := &jsonCapQuery{}
|
||||
h.applyFilter(q2, common.FilterOption{
|
||||
Column: "data.city", Operator: "eq", Value: "LA",
|
||||
}, "public.things", false, "AND", nil)
|
||||
if c := q2.only(t); c.query == `("things"."data" #>> ?::text[]) = ?` {
|
||||
t.Fatalf("dotted shorthand should not resolve without a model: %+v", c)
|
||||
}
|
||||
})
|
||||
}
|
||||
+59
-17
@@ -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)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -336,21 +336,21 @@ type (
|
||||
SqlUUID = SqlNull[uuid.UUID]
|
||||
)
|
||||
|
||||
// SqlTimeStamp - Timestamp with custom formatting (YYYY-MM-DDTHH:MM:SS).
|
||||
// SqlTimeStamp - Timestamp serialized as RFC3339 with timezone offset.
|
||||
type SqlTimeStamp struct{ SqlNull[time.Time] }
|
||||
|
||||
func (t SqlTimeStamp) MarshalJSON() ([]byte, error) {
|
||||
if !t.Valid || t.Val.IsZero() || t.Val.Before(time.Date(0002, 1, 1, 0, 0, 0, 0, time.UTC)) {
|
||||
return []byte("null"), nil
|
||||
}
|
||||
return []byte(fmt.Sprintf(`"%s"`, t.Val.Format("2006-01-02T15:04:05"))), nil
|
||||
return []byte(fmt.Sprintf(`"%s"`, t.Val.Format(time.RFC3339))), nil
|
||||
}
|
||||
|
||||
func (t *SqlTimeStamp) UnmarshalJSON(b []byte) error {
|
||||
if err := t.SqlNull.UnmarshalJSON(b); err != nil {
|
||||
return err
|
||||
}
|
||||
if t.Valid && (t.Val.IsZero() || t.Val.Format("2006-01-02T15:04:05") == "0001-01-01T00:00:00") {
|
||||
if t.Valid && (t.Val.IsZero() || t.Val.Year() <= 1) {
|
||||
t.Valid = false
|
||||
}
|
||||
return nil
|
||||
@@ -360,7 +360,7 @@ func (t SqlTimeStamp) Value() (driver.Value, error) {
|
||||
if !t.Valid || t.Val.IsZero() || t.Val.Before(time.Date(0002, 1, 1, 0, 0, 0, 0, time.UTC)) {
|
||||
return nil, nil
|
||||
}
|
||||
return t.Val.Format("2006-01-02T15:04:05"), nil
|
||||
return t.Val.Format(time.RFC3339), nil
|
||||
}
|
||||
|
||||
func SqlTimeStampNow() SqlTimeStamp {
|
||||
@@ -425,9 +425,7 @@ func (t SqlTime) MarshalJSON() ([]byte, error) {
|
||||
return []byte("null"), nil
|
||||
}
|
||||
s := t.Val.Format("15:04:05")
|
||||
if s == "00:00:00" {
|
||||
return []byte("null"), nil
|
||||
}
|
||||
|
||||
return []byte(fmt.Sprintf(`"%s"`, s)), nil
|
||||
}
|
||||
|
||||
@@ -435,9 +433,7 @@ func (t *SqlTime) UnmarshalJSON(b []byte) error {
|
||||
if err := t.SqlNull.UnmarshalJSON(b); err != nil {
|
||||
return err
|
||||
}
|
||||
if t.Valid && t.Val.Format("15:04:05") == "00:00:00" {
|
||||
t.Valid = false
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
@@ -178,7 +178,7 @@ func TestSqlTimeStamp_JSON(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("Marshal failed: %v", err)
|
||||
}
|
||||
expected := `"2024-01-15T10:30:45"`
|
||||
expected := `"2024-01-15T10:30:45Z"`
|
||||
if string(data) != expected {
|
||||
t.Errorf("expected %s, got %s", expected, string(data))
|
||||
}
|
||||
@@ -920,6 +920,411 @@ func TestSqlString_RoundTrip(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// TestSqlBool_Scan tests SqlBool Scan from various input types.
|
||||
func TestSqlBool_Scan(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
input interface{}
|
||||
expected bool
|
||||
valid bool
|
||||
}{
|
||||
{"bool true", true, true, true},
|
||||
{"bool false", false, false, true},
|
||||
{"string true", "true", true, true},
|
||||
{"string 1", "1", true, true},
|
||||
{"int64 1 fallback", int64(1), true, true},
|
||||
{"int64 0 fallback", int64(0), false, true},
|
||||
{"nil", nil, false, false},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
var b SqlBool
|
||||
if err := b.Scan(tt.input); err != nil {
|
||||
t.Fatalf("Scan failed: %v", err)
|
||||
}
|
||||
if b.Valid != tt.valid {
|
||||
t.Errorf("expected valid=%v, got valid=%v", tt.valid, b.Valid)
|
||||
}
|
||||
if tt.valid && b.Val != tt.expected {
|
||||
t.Errorf("expected %v, got %v", tt.expected, b.Val)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSqlBool_Value(t *testing.T) {
|
||||
b := NewSqlBool(true)
|
||||
val, err := b.Value()
|
||||
if err != nil {
|
||||
t.Fatalf("Value failed: %v", err)
|
||||
}
|
||||
if val != true {
|
||||
t.Errorf("expected true, got %v", val)
|
||||
}
|
||||
|
||||
b2 := SqlBool{Valid: false}
|
||||
val2, err := b2.Value()
|
||||
if err != nil {
|
||||
t.Fatalf("Value failed: %v", err)
|
||||
}
|
||||
if val2 != nil {
|
||||
t.Errorf("expected nil, got %v", val2)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSqlBool_JSON(t *testing.T) {
|
||||
b := NewSqlBool(true)
|
||||
data, err := json.Marshal(b)
|
||||
if err != nil {
|
||||
t.Fatalf("Marshal failed: %v", err)
|
||||
}
|
||||
if string(data) != "true" {
|
||||
t.Errorf("expected true, got %s", string(data))
|
||||
}
|
||||
|
||||
var b2 SqlBool
|
||||
if err := json.Unmarshal([]byte("false"), &b2); err != nil {
|
||||
t.Fatalf("Unmarshal failed: %v", err)
|
||||
}
|
||||
if !b2.Valid || b2.Val != false {
|
||||
t.Errorf("expected valid=true val=false, got valid=%v val=%v", b2.Valid, b2.Val)
|
||||
}
|
||||
|
||||
var b3 SqlBool
|
||||
if err := json.Unmarshal([]byte("null"), &b3); err != nil {
|
||||
t.Fatalf("Unmarshal null failed: %v", err)
|
||||
}
|
||||
if b3.Valid {
|
||||
t.Error("expected invalid after unmarshaling null")
|
||||
}
|
||||
}
|
||||
|
||||
// TestSqlNull_FromString_EdgeCases tests FromString edge cases shared by all SqlNull instantiations.
|
||||
func TestSqlNull_FromString_EdgeCases(t *testing.T) {
|
||||
t.Run("empty string is null", func(t *testing.T) {
|
||||
var n SqlInt64
|
||||
if err := n.FromString(""); err != nil {
|
||||
t.Fatalf("FromString failed: %v", err)
|
||||
}
|
||||
if n.Valid {
|
||||
t.Error("expected invalid for empty string")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("NULL case-insensitive", func(t *testing.T) {
|
||||
var n SqlString
|
||||
if err := n.FromString("NuLL"); err != nil {
|
||||
t.Fatalf("FromString failed: %v", err)
|
||||
}
|
||||
if n.Valid {
|
||||
t.Error("expected invalid for 'NuLL'")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("whitespace trimmed", func(t *testing.T) {
|
||||
var n SqlInt64
|
||||
if err := n.FromString(" 42 "); err != nil {
|
||||
t.Fatalf("FromString failed: %v", err)
|
||||
}
|
||||
if !n.Valid || n.Val != 42 {
|
||||
t.Errorf("expected valid=true val=42, got valid=%v val=%v", n.Valid, n.Val)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("invalid int string stays invalid", func(t *testing.T) {
|
||||
var n SqlInt64
|
||||
if err := n.FromString("not-a-number"); err != nil {
|
||||
t.Fatalf("FromString failed: %v", err)
|
||||
}
|
||||
if n.Valid {
|
||||
t.Error("expected invalid for non-numeric string")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("float string truncated into int type", func(t *testing.T) {
|
||||
var n SqlInt64
|
||||
if err := n.FromString("3.7"); err != nil {
|
||||
t.Fatalf("FromString failed: %v", err)
|
||||
}
|
||||
if !n.Valid || n.Val != 3 {
|
||||
t.Errorf("expected valid=true val=3, got valid=%v val=%v", n.Valid, n.Val)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("invalid bool string stays invalid", func(t *testing.T) {
|
||||
var n SqlBool
|
||||
if err := n.FromString("maybe"); err != nil {
|
||||
t.Fatalf("FromString failed: %v", err)
|
||||
}
|
||||
if n.Valid {
|
||||
t.Error("expected invalid for non-bool string")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// TestSqlNull_String tests the String() stringer fallback.
|
||||
func TestSqlNull_String(t *testing.T) {
|
||||
t.Run("invalid returns empty", func(t *testing.T) {
|
||||
n := SqlInt64{Valid: false}
|
||||
if n.String() != "" {
|
||||
t.Errorf("expected empty string, got %q", n.String())
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("stringer type delegates", func(t *testing.T) {
|
||||
u := uuid.New()
|
||||
n := NewSqlUUID(u)
|
||||
if n.String() != u.String() {
|
||||
t.Errorf("expected %s, got %s", u.String(), n.String())
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("non-stringer falls back to fmt", func(t *testing.T) {
|
||||
n := NewSqlInt64(42)
|
||||
if n.String() != "42" {
|
||||
t.Errorf("expected 42, got %s", n.String())
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// TestNewSql_Generic tests the generic NewSql constructor.
|
||||
func TestNewSql_Generic(t *testing.T) {
|
||||
t.Run("exact type match", func(t *testing.T) {
|
||||
n := NewSql[int64](int64(5))
|
||||
if !n.Valid || n.Val != 5 {
|
||||
t.Errorf("expected valid=true val=5, got valid=%v val=%v", n.Valid, n.Val)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("nil value", func(t *testing.T) {
|
||||
n := NewSql[int64](nil)
|
||||
if n.Valid {
|
||||
t.Error("expected invalid for nil")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("from another SqlNull", func(t *testing.T) {
|
||||
src := SqlNull[int64]{Val: 9, Valid: true}
|
||||
n := NewSql[int64](src)
|
||||
if !n.Valid || n.Val != 9 {
|
||||
t.Errorf("expected valid=true val=9, got valid=%v val=%v", n.Valid, n.Val)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("string conversion fallback", func(t *testing.T) {
|
||||
n := NewSql[string](42)
|
||||
if !n.Valid || n.Val != "42" {
|
||||
t.Errorf("expected valid=true val=42, got valid=%v val=%q", n.Valid, n.Val)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// TestSqlNull_Int64_Conversions tests Int64() across differently-typed SqlNull values.
|
||||
func TestSqlNull_Int64_Conversions(t *testing.T) {
|
||||
if v := (SqlNull[string]{Val: "42", Valid: true}).Int64(); v != 42 {
|
||||
t.Errorf("expected 42, got %d", v)
|
||||
}
|
||||
if v := (SqlNull[bool]{Val: true, Valid: true}).Int64(); v != 1 {
|
||||
t.Errorf("expected 1, got %d", v)
|
||||
}
|
||||
if v := (SqlNull[bool]{Val: false, Valid: true}).Int64(); v != 0 {
|
||||
t.Errorf("expected 0, got %d", v)
|
||||
}
|
||||
if v := (SqlNull[float64]{Val: 3.9, Valid: true}).Int64(); v != 3 {
|
||||
t.Errorf("expected 3, got %d", v)
|
||||
}
|
||||
if v := (SqlNull[int64]{Valid: false}).Int64(); v != 0 {
|
||||
t.Errorf("expected 0 for invalid, got %d", v)
|
||||
}
|
||||
}
|
||||
|
||||
// TestSqlNull_Float64_Conversions tests Float64() across differently-typed SqlNull values.
|
||||
func TestSqlNull_Float64_Conversions(t *testing.T) {
|
||||
if v := (SqlNull[string]{Val: "3.14", Valid: true}).Float64(); v != 3.14 {
|
||||
t.Errorf("expected 3.14, got %v", v)
|
||||
}
|
||||
if v := (SqlNull[int64]{Val: 10, Valid: true}).Float64(); v != 10.0 {
|
||||
t.Errorf("expected 10.0, got %v", v)
|
||||
}
|
||||
if v := (SqlNull[float64]{Valid: false}).Float64(); v != 0.0 {
|
||||
t.Errorf("expected 0.0 for invalid, got %v", v)
|
||||
}
|
||||
}
|
||||
|
||||
// TestSqlNull_Bool_Conversions tests Bool() across differently-typed SqlNull values.
|
||||
func TestSqlNull_Bool_Conversions(t *testing.T) {
|
||||
if v := (SqlNull[string]{Val: "YES", Valid: true}).Bool(); v != true {
|
||||
t.Error("expected true for 'YES'")
|
||||
}
|
||||
if v := (SqlNull[string]{Val: "no", Valid: true}).Bool(); v != false {
|
||||
t.Error("expected false for 'no'")
|
||||
}
|
||||
if v := (SqlNull[int]{Val: 1, Valid: true}).Bool(); v != true {
|
||||
t.Error("expected true for int 1")
|
||||
}
|
||||
if v := (SqlNull[int]{Val: 0, Valid: true}).Bool(); v != false {
|
||||
t.Error("expected false for int 0")
|
||||
}
|
||||
if v := (SqlNull[bool]{Valid: false}).Bool(); v != false {
|
||||
t.Error("expected false for invalid")
|
||||
}
|
||||
}
|
||||
|
||||
// TestSqlNull_Time_NonTimeType verifies Time() returns zero value when T is not time.Time.
|
||||
func TestSqlNull_Time_NonTimeType(t *testing.T) {
|
||||
n := SqlNull[string]{Val: "2024-01-15", Valid: true}
|
||||
if !n.Time().IsZero() {
|
||||
t.Error("expected zero time for non-time.Time SqlNull")
|
||||
}
|
||||
}
|
||||
|
||||
// TestSqlNull_UUID_NonUUIDType verifies UUID() returns uuid.Nil when T is not uuid.UUID.
|
||||
func TestSqlNull_UUID_NonUUIDType(t *testing.T) {
|
||||
n := SqlNull[string]{Val: "not-a-uuid", Valid: true}
|
||||
if n.UUID() != uuid.Nil {
|
||||
t.Error("expected uuid.Nil for non-uuid.UUID SqlNull")
|
||||
}
|
||||
}
|
||||
|
||||
// TestSqlTime_Midnight verifies midnight times are serialized as "00:00:00", not null.
|
||||
func TestSqlTime_Midnight(t *testing.T) {
|
||||
midnight := time.Date(2024, 1, 1, 0, 0, 0, 0, time.UTC)
|
||||
tm := NewSqlTime(midnight)
|
||||
|
||||
data, err := json.Marshal(tm)
|
||||
if err != nil {
|
||||
t.Fatalf("Marshal failed: %v", err)
|
||||
}
|
||||
if string(data) != `"00:00:00"` {
|
||||
t.Errorf("expected \"00:00:00\", got %s", string(data))
|
||||
}
|
||||
|
||||
val, err := tm.Value()
|
||||
if err != nil {
|
||||
t.Fatalf("Value failed: %v", err)
|
||||
}
|
||||
if val != "00:00:00" {
|
||||
t.Errorf("expected 00:00:00, got %v", val)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSqlTime_Value_Invalid(t *testing.T) {
|
||||
tm := SqlTime{}
|
||||
val, err := tm.Value()
|
||||
if err != nil {
|
||||
t.Fatalf("Value failed: %v", err)
|
||||
}
|
||||
if val != nil {
|
||||
t.Errorf("expected nil, got %v", val)
|
||||
}
|
||||
}
|
||||
|
||||
// TestSqlDate_ZeroValueString verifies String() blanks out sentinel zero dates.
|
||||
func TestSqlDate_ZeroValueString(t *testing.T) {
|
||||
d := SqlDate{SqlNull: SqlNull[time.Time]{Val: time.Time{}, Valid: true}}
|
||||
if d.String() != "" {
|
||||
t.Errorf("expected empty string for zero date, got %q", d.String())
|
||||
}
|
||||
|
||||
sentinel := time.Date(1800, 12, 31, 0, 0, 0, 0, time.UTC)
|
||||
d2 := SqlDate{SqlNull: SqlNull[time.Time]{Val: sentinel, Valid: true}}
|
||||
if d2.String() != "" {
|
||||
t.Errorf("expected empty string for 1800-12-31 sentinel, got %q", d2.String())
|
||||
}
|
||||
}
|
||||
|
||||
// TestSqlTimeStamp_Value tests driver.Valuer for SqlTimeStamp, including the pre-year-2 cutoff.
|
||||
func TestSqlTimeStamp_Value(t *testing.T) {
|
||||
t.Run("valid recent timestamp", func(t *testing.T) {
|
||||
ts := NewSqlTimeStamp(time.Date(2024, 1, 15, 10, 30, 0, 0, time.UTC))
|
||||
val, err := ts.Value()
|
||||
if err != nil {
|
||||
t.Fatalf("Value failed: %v", err)
|
||||
}
|
||||
if val != "2024-01-15T10:30:00Z" {
|
||||
t.Errorf("expected 2024-01-15T10:30:00Z, got %v", val)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("year 1 is treated as null", func(t *testing.T) {
|
||||
ts := NewSqlTimeStamp(time.Date(1, 1, 1, 0, 0, 0, 0, time.UTC))
|
||||
val, err := ts.Value()
|
||||
if err != nil {
|
||||
t.Fatalf("Value failed: %v", err)
|
||||
}
|
||||
if val != nil {
|
||||
t.Errorf("expected nil, got %v", val)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("invalid is null", func(t *testing.T) {
|
||||
ts := SqlTimeStamp{}
|
||||
val, err := ts.Value()
|
||||
if err != nil {
|
||||
t.Fatalf("Value failed: %v", err)
|
||||
}
|
||||
if val != nil {
|
||||
t.Errorf("expected nil, got %v", val)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestSqlTimeStamp_UnmarshalJSON_YearOneInvalid(t *testing.T) {
|
||||
var ts SqlTimeStamp
|
||||
if err := json.Unmarshal([]byte(`"0001-01-01T00:00:00Z"`), &ts); err != nil {
|
||||
t.Fatalf("Unmarshal failed: %v", err)
|
||||
}
|
||||
if ts.Valid {
|
||||
t.Error("expected invalid for year 0001 timestamp")
|
||||
}
|
||||
}
|
||||
|
||||
// TestTryParseDT tests the internal multi-format date/time parser.
|
||||
func TestTryParseDT(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
input string
|
||||
}{
|
||||
{"RFC3339", "2024-01-15T10:30:00Z"},
|
||||
{"date only", "2024-01-15"},
|
||||
{"datetime no tz", "2024-01-15T10:30:00"},
|
||||
{"space separated", "2024-01-15 10:30:00"},
|
||||
{"UK date slash", "15/01/2024"},
|
||||
{"UK date dash", "15-01-2024"},
|
||||
{"time only", "10:30:00"},
|
||||
{"short time", "10:30"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
tm, err := tryParseDT(tt.input)
|
||||
if err != nil {
|
||||
t.Fatalf("tryParseDT failed for %q: %v", tt.input, err)
|
||||
}
|
||||
if tm.IsZero() {
|
||||
t.Errorf("expected non-zero time for %q", tt.input)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
t.Run("invalid format", func(t *testing.T) {
|
||||
_, err := tryParseDT("not a date at all")
|
||||
if err == nil {
|
||||
t.Error("expected error for unparseable string")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// TestToJSONDT tests RFC3339 formatting helper.
|
||||
func TestToJSONDT(t *testing.T) {
|
||||
dt := time.Date(2024, 1, 15, 10, 30, 0, 0, time.UTC)
|
||||
expected := dt.Format(time.RFC3339)
|
||||
if got := ToJSONDT(dt); got != expected {
|
||||
t.Errorf("expected %s, got %s", expected, got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestSqlByteArray_Base64_RoundTrip tests complete round-trip: Go -> JSON -> Go -> SQL -> Go
|
||||
func TestSqlByteArray_Base64_RoundTrip(t *testing.T) {
|
||||
original := []byte{0x48, 0x65, 0x6C, 0x6C, 0x6F, 0x20, 0xFF, 0xFE} // "Hello " + binary data
|
||||
|
||||
@@ -91,3 +91,33 @@ func IsVectorType(t reflect.Type) bool {
|
||||
n, ok := SQLTypeName(t)
|
||||
return ok && (n == "vector" || n == "halfvec" || n == "sparsevec")
|
||||
}
|
||||
|
||||
// IsJSONType reports whether t is a spectypes JSON/JSONB wrapper.
|
||||
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()
|
||||
}
|
||||
|
||||
@@ -51,6 +51,20 @@ func TestIsSpatialType(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsJSONType(t *testing.T) {
|
||||
if !IsJSONType(reflect.TypeOf(SqlJSONB{})) {
|
||||
t.Error("SqlJSONB should be a JSON type")
|
||||
}
|
||||
if !IsJSONType(reflect.TypeOf(&SqlJSONB{})) {
|
||||
t.Error("*SqlJSONB should be a JSON type (pointer unwrapped)")
|
||||
}
|
||||
for _, v := range []any{SqlGeometry{}, SqlVector{}, SqlString{}, SqlStringArray{}, ""} {
|
||||
if IsJSONType(reflect.TypeOf(v)) {
|
||||
t.Errorf("%T should not be a JSON type", v)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsVectorType(t *testing.T) {
|
||||
for _, v := range []any{SqlVector{}, SqlHalfVector{}, SqlSparseVector{}} {
|
||||
if !IsVectorType(reflect.TypeOf(v)) {
|
||||
|
||||
@@ -564,7 +564,7 @@ func (h *Handler) readByID(hookCtx *HookContext) (interface{}, error) {
|
||||
|
||||
// Apply columns
|
||||
if hookCtx.Options != nil && len(hookCtx.Options.Columns) > 0 {
|
||||
query = query.Column(hookCtx.Options.Columns...)
|
||||
query = common.ApplySelectColumns(query, hookCtx.Model, "", hookCtx.Options.Columns)
|
||||
}
|
||||
|
||||
// Apply preloads (simplified for now)
|
||||
@@ -606,7 +606,7 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata
|
||||
// Apply options (simplified implementation)
|
||||
if hookCtx.Options != nil {
|
||||
// Apply filters with OR grouping support
|
||||
query = h.applyFilters(query, hookCtx.Options.Filters)
|
||||
query = h.applyFilters(query, hookCtx.Options.Filters, hookCtx.Model)
|
||||
|
||||
// Apply sorting
|
||||
for _, sort := range hookCtx.Options.Sort {
|
||||
@@ -614,6 +614,10 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata
|
||||
if sort.Direction == "desc" {
|
||||
direction = "DESC"
|
||||
}
|
||||
if expr, jargs, _, ok := common.ResolveJSONColumnExpr(hookCtx.Model, "", sort.Column); ok {
|
||||
query = query.OrderExpr(fmt.Sprintf("%s %s", expr, direction), jargs...)
|
||||
continue
|
||||
}
|
||||
query = query.Order(fmt.Sprintf("%s %s", sort.Column, direction))
|
||||
}
|
||||
|
||||
@@ -632,7 +636,7 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata
|
||||
|
||||
// Apply columns
|
||||
if len(hookCtx.Options.Columns) > 0 {
|
||||
query = query.Column(hookCtx.Options.Columns...)
|
||||
query = common.ApplySelectColumns(query, hookCtx.Model, "", hookCtx.Options.Columns)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -665,7 +669,7 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata
|
||||
countQuery := h.db.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName)
|
||||
if hookCtx.Options != nil {
|
||||
for _, filter := range hookCtx.Options.Filters {
|
||||
cond, args := h.buildFilterCondition(filter)
|
||||
cond, args := h.buildFilterCondition(filter, hookCtx.Model)
|
||||
if cond != "" {
|
||||
countQuery = countQuery.Where(cond, args...)
|
||||
}
|
||||
@@ -776,7 +780,7 @@ func (h *Handler) getMetadata(schema, entity string, model interface{}) map[stri
|
||||
// getOperatorSQL converts filter operator to SQL operator
|
||||
// applyFilters applies all filters with proper grouping for OR logic
|
||||
// Groups consecutive OR filters together to ensure proper query precedence
|
||||
func (h *Handler) applyFilters(query common.SelectQuery, filters []common.FilterOption) common.SelectQuery {
|
||||
func (h *Handler) applyFilters(query common.SelectQuery, filters []common.FilterOption, model interface{}) common.SelectQuery {
|
||||
if len(filters) == 0 {
|
||||
return query
|
||||
}
|
||||
@@ -796,11 +800,11 @@ func (h *Handler) applyFilters(query common.SelectQuery, filters []common.Filter
|
||||
}
|
||||
|
||||
// Apply the OR group as a single grouped WHERE clause
|
||||
query = h.applyFilterGroup(query, orGroup)
|
||||
query = h.applyFilterGroup(query, orGroup, model)
|
||||
i = j
|
||||
} else {
|
||||
// Single filter with AND logic (or first filter)
|
||||
condition, args := h.buildFilterCondition(filters[i])
|
||||
condition, args := h.buildFilterCondition(filters[i], model)
|
||||
if condition != "" {
|
||||
query = query.Where(condition, args...)
|
||||
}
|
||||
@@ -813,7 +817,7 @@ func (h *Handler) applyFilters(query common.SelectQuery, filters []common.Filter
|
||||
|
||||
// applyFilterGroup applies a group of filters that should be OR'd together
|
||||
// Always wraps them in parentheses and applies as a single WHERE clause
|
||||
func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.FilterOption) common.SelectQuery {
|
||||
func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.FilterOption, model interface{}) common.SelectQuery {
|
||||
if len(filters) == 0 {
|
||||
return query
|
||||
}
|
||||
@@ -823,7 +827,7 @@ func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.Fi
|
||||
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...)
|
||||
@@ -844,8 +848,14 @@ func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.Fi
|
||||
return query.Where(groupedCondition, args...)
|
||||
}
|
||||
|
||||
// buildFilterCondition builds a filter condition and returns it with args
|
||||
func (h *Handler) buildFilterCondition(filter common.FilterOption) (conditionString string, conditionArgs []interface{}) {
|
||||
// buildFilterCondition builds a filter condition and returns it with args.
|
||||
// model, when non-nil, lets JSON sub-field references (data->>'x', data#>>'{a,b}',
|
||||
// or the dotted data.x shorthand for a JSON column) resolve to a safe,
|
||||
// parameterised expression before the ordinary operator handling below.
|
||||
func (h *Handler) buildFilterCondition(filter common.FilterOption, model interface{}) (conditionString string, conditionArgs []interface{}) {
|
||||
if cond, jargs, ok := common.BuildJSONFilterCondition(model, "", filter.Column, filter.Operator, filter.Value); ok {
|
||||
return cond, jargs
|
||||
}
|
||||
if strings.EqualFold(filter.Operator, "in") {
|
||||
cond, args := common.BuildInCondition(filter.Column, filter.Value)
|
||||
return cond, args
|
||||
@@ -853,6 +863,11 @@ func (h *Handler) buildFilterCondition(filter common.FilterOption) (conditionStr
|
||||
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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
Vendored
+1
-1
File diff suppressed because one or more lines are too long
Vendored
+5
-366
@@ -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
|
||||
Vendored
+426
-463
@@ -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
@@ -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"
|
||||
|
||||
Generated
+1283
-1293
File diff suppressed because it is too large
Load Diff
@@ -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),
|
||||
});
|
||||
}
|
||||
|
||||
@@ -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]);
|
||||
}
|
||||
@@ -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>;
|
||||
}
|
||||
|
||||
@@ -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),
|
||||
});
|
||||
}
|
||||
|
||||
@@ -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>> {
|
||||
|
||||
@@ -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'],
|
||||
},
|
||||
},
|
||||
});
|
||||
|
||||
Reference in New Issue
Block a user