From 6098526faa2bdd95775952f197b429ed316516c3 Mon Sep 17 00:00:00 2001 From: Nitin Kumar Date: Sun, 4 Oct 2026 17:55:12 +0530 Subject: [PATCH 1/3] feat(aws-apigateway): method/integration responses, mapping templates, MOCK (A2) --- docs/coverage/README.md | 2 +- docs/coverage/aws/README.md | 2 +- docs/coverage/aws/apigateway.md | 10 +- docs/coverage/coverage.json | 28 + internal/jsonpath/jsonpath.go | 168 +++ internal/jsonpath/jsonpath_test.go | 61 + internal/vtl/ast.go | 110 ++ internal/vtl/doc.go | 11 + internal/vtl/edge_test.go | 129 ++ internal/vtl/errors.go | 15 + internal/vtl/eval.go | 724 +++++++++++ internal/vtl/methods.go | 430 +++++++ internal/vtl/parse.go | 1109 +++++++++++++++++ internal/vtl/value.go | 345 +++++ internal/vtl/vtl_test.go | 186 +++ providers/aws/apigateway/dataplane.go | 10 +- providers/aws/apigateway/mapping.go | 356 ++++++ .../aws/apigateway/mapping_validation.go | 159 +++ providers/aws/apigateway/methods.go | 31 +- providers/aws/apigateway/mock_integration.go | 347 ++++++ .../aws/apigateway/mock_integration_test.go | 267 ++++ providers/aws/apigateway/patch.go | 84 +- providers/aws/apigateway/resources.go | 8 +- providers/aws/apigateway/responses.go | 427 +++++++ providers/aws/sfn/asl/jsonpath.go | 137 +- server/aws/apigateway/handler.go | 15 +- .../apigateway/mock_integration_e2e_test.go | 211 ++++ server/aws/apigateway/responses.go | 103 ++ server/aws/apigateway/types.go | 126 +- services/apigateway/driver/driver.go | 109 +- 30 files changed, 5525 insertions(+), 195 deletions(-) create mode 100644 internal/jsonpath/jsonpath.go create mode 100644 internal/jsonpath/jsonpath_test.go create mode 100644 internal/vtl/ast.go create mode 100644 internal/vtl/doc.go create mode 100644 internal/vtl/edge_test.go create mode 100644 internal/vtl/errors.go create mode 100644 internal/vtl/eval.go create mode 100644 internal/vtl/methods.go create mode 100644 internal/vtl/parse.go create mode 100644 internal/vtl/value.go create mode 100644 internal/vtl/vtl_test.go create mode 100644 providers/aws/apigateway/mapping.go create mode 100644 providers/aws/apigateway/mapping_validation.go create mode 100644 providers/aws/apigateway/mock_integration.go create mode 100644 providers/aws/apigateway/mock_integration_test.go create mode 100644 providers/aws/apigateway/responses.go create mode 100644 server/aws/apigateway/mock_integration_e2e_test.go create mode 100644 server/aws/apigateway/responses.go diff --git a/docs/coverage/README.md b/docs/coverage/README.md index b9a759762..bf8cbf459 100644 --- a/docs/coverage/README.md +++ b/docs/coverage/README.md @@ -15,7 +15,7 @@ code does not implement. Machine-readable: [`coverage.json`](./coverage.json). | `acm` | [ACM](./aws/acm.md) | - | - | - | 17 | | `aks` | - | [AKS](./azure/aks.md) | - | - | 19 | | `aoss` | [AOSS](./aws/aoss.md) | - | - | - | 18 | -| `apigateway` | [APIGateway](./aws/apigateway.md) | - | - | - | 50 | +| `apigateway` | [APIGateway](./aws/apigateway.md) | - | - | - | 58 | | `apigatewaygcp` | - | - | [APIGateway](./gcp/apigateway.md) | - | 16 | | `apigatewayv2` | [APIGatewayV2](./aws/apigatewayv2.md) | - | - | - | 28 | | `apimanagement` | - | [APIManagement](./azure/apimanagement.md) | - | - | 30 | diff --git a/docs/coverage/aws/README.md b/docs/coverage/aws/README.md index 216629258..fe7b8ef90 100644 --- a/docs/coverage/aws/README.md +++ b/docs/coverage/aws/README.md @@ -7,7 +7,7 @@ Services cloudemu emulates for AWS, by native name. Back to the [cross-provider | --- | --- | --- | | [ACM](./acm.md) | `acm` | 17 | | [AOSS](./aoss.md) | `aoss` | 18 | -| [APIGateway](./apigateway.md) | `apigateway` | 50 | +| [APIGateway](./apigateway.md) | `apigateway` | 58 | | [APIGatewayV2](./apigatewayv2.md) | `apigatewayv2` | 28 | | [APS](./aps.md) | `aps` | 21 | | [AppFlow](./appflow.md) | `appflow` | 14 | diff --git a/docs/coverage/aws/apigateway.md b/docs/coverage/aws/apigateway.md index 98b04c86f..6a037e5ce 100644 --- a/docs/coverage/aws/apigateway.md +++ b/docs/coverage/aws/apigateway.md @@ -3,7 +3,7 @@ AWS's `apigateway` service · portable interface `driver.APIGateway` · [AWS index](./README.md) -## Operations (50) +## Operations (58) | Operation | Description | | --- | --- | @@ -18,7 +18,9 @@ AWS's `apigateway` service · portable interface `driver.APIGateway` · [AWS ind | `DeleteDocumentationPart` | | | `DeleteDocumentationVersion` | DeleteDocumentationVersion removes a version. It fails while a stage | | `DeleteIntegration` | | +| `DeleteIntegrationResponse` | | | `DeleteMethod` | | +| `DeleteMethodResponse` | | | `DeleteResource` | DeleteResource removes a resource and its whole descendant subtree, as | | `DeleteRestAPI` | | | `DeleteStage` | | @@ -33,7 +35,9 @@ AWS's `apigateway` service · portable interface `driver.APIGateway` · [AWS ind | `GetDocumentationVersion` | | | `GetDocumentationVersions` | | | `GetIntegration` | | +| `GetIntegrationResponse` | | | `GetMethod` | | +| `GetMethodResponse` | | | `GetResource` | | | `GetResources` | | | `GetRestAPI` | | @@ -44,7 +48,9 @@ AWS's `apigateway` service · portable interface `driver.APIGateway` · [AWS ind | `ImportDocumentationParts` | | | `InvokeRoute` | InvokeRoute routes req through the tree its stage's deployment captured. | | `PutIntegration` | | +| `PutIntegrationResponse` | PutIntegrationResponse creates or replaces an integration response. The | | `PutMethod` | | +| `PutMethodResponse` | PutMethodResponse declares a status code on a method. It fails when the | | `TagResource` | TagResource, UntagResource and GetTags manage tags on a REST API or client | | `UntagResource` | | | `UpdateAccount` | UpdateAccount applies a patchOperations document (/cloudwatchRoleArn | @@ -53,7 +59,9 @@ AWS's `apigateway` service · portable interface `driver.APIGateway` · [AWS ind | `UpdateDocumentationPart` | UpdateDocumentationPart applies a patchOperations document (only | | `UpdateDocumentationVersion` | UpdateDocumentationVersion applies a patchOperations document (only | | `UpdateIntegration` | UpdateIntegration applies a patchOperations document to an integration. | +| `UpdateIntegrationResponse` | UpdateIntegrationResponse applies a patchOperations document | | `UpdateMethod` | UpdateMethod applies a patchOperations document to a method. | +| `UpdateMethodResponse` | UpdateMethodResponse applies a patchOperations document | | `UpdateResource` | UpdateResource applies a patchOperations document to a resource (rename via | | `UpdateRestAPI` | UpdateRestAPI applies a patchOperations document to a REST API and returns | | `UpdateStage` | UpdateStage applies a patchOperations document to a stage (/description, | diff --git a/docs/coverage/coverage.json b/docs/coverage/coverage.json index 8c3fc5a3a..01e529e55 100644 --- a/docs/coverage/coverage.json +++ b/docs/coverage/coverage.json @@ -306,9 +306,15 @@ { "name": "DeleteIntegration" }, + { + "name": "DeleteIntegrationResponse" + }, { "name": "DeleteMethod" }, + { + "name": "DeleteMethodResponse" + }, { "name": "DeleteResource", "doc": "DeleteResource removes a resource and its whole descendant subtree, as" @@ -353,9 +359,15 @@ { "name": "GetIntegration" }, + { + "name": "GetIntegrationResponse" + }, { "name": "GetMethod" }, + { + "name": "GetMethodResponse" + }, { "name": "GetResource" }, @@ -387,9 +399,17 @@ { "name": "PutIntegration" }, + { + "name": "PutIntegrationResponse", + "doc": "PutIntegrationResponse creates or replaces an integration response. The" + }, { "name": "PutMethod" }, + { + "name": "PutMethodResponse", + "doc": "PutMethodResponse declares a status code on a method. It fails when the" + }, { "name": "TagResource", "doc": "TagResource, UntagResource and GetTags manage tags on a REST API or client" @@ -421,10 +441,18 @@ "name": "UpdateIntegration", "doc": "UpdateIntegration applies a patchOperations document to an integration." }, + { + "name": "UpdateIntegrationResponse", + "doc": "UpdateIntegrationResponse applies a patchOperations document" + }, { "name": "UpdateMethod", "doc": "UpdateMethod applies a patchOperations document to a method." }, + { + "name": "UpdateMethodResponse", + "doc": "UpdateMethodResponse applies a patchOperations document" + }, { "name": "UpdateResource", "doc": "UpdateResource applies a patchOperations document to a resource (rename via" diff --git a/internal/jsonpath/jsonpath.go b/internal/jsonpath/jsonpath.go new file mode 100644 index 000000000..f80a56fbb --- /dev/null +++ b/internal/jsonpath/jsonpath.go @@ -0,0 +1,168 @@ +// Package jsonpath evaluates the reference subset of JSONPath shared by the +// Step Functions ASL interpreter and the API Gateway mapping templates: "$", +// "$.a.b", "$[0]", "$['a']" and "$.a.b[2]". Filters, wildcards and recursive +// descent are rejected with an error, so an unsupported path fails loudly +// instead of returning a wrong result. +// +// Values are plain decoded JSON (map[string]any, []any) or any type that +// implements Object or Array, which lets callers keep an order-preserving +// document model. +package jsonpath + +import ( + "fmt" + "strconv" + "strings" +) + +// Object is a JSON object view a path can step into by field name. +type Object interface { + Lookup(key string) (any, bool) +} + +// Array is a JSON array view a path can step into by index. +type Array interface { + Index(i int) (any, bool) +} + +// Error reports a malformed or unsupported path. +type Error struct { + Msg string +} + +func (e *Error) Error() string { return e.Msg } + +func errorf(format string, args ...any) error { + return &Error{Msg: fmt.Sprintf(format, args...)} +} + +// Eval evaluates path against root and returns the selected value and whether +// it was present. +func Eval(path string, root any) (value any, present bool, err error) { + if path == "" || path[0] != '$' { + return nil, false, errorf("invalid JSONPath %q: must start with '$'", path) + } + + if strings.ContainsAny(path, "*?@") || strings.Contains(path, "..") { + return nil, false, errorf("JSONPath %q uses unsupported syntax (filters/wildcards/recursive descent)", path) + } + + if path == "$" { + return root, true, nil + } + + toks, err := tokenize(path[1:]) + if err != nil { + return nil, false, err + } + + cur := root + + for _, t := range toks { + next, ok := t.apply(cur) + if !ok { + return nil, false, nil + } + + cur = next + } + + return cur, true, nil +} + +// token is one selection step: a field name or an array index. +type token struct { + field string + index int + isIndex bool +} + +func (t token) apply(cur any) (any, bool) { + if t.isIndex { + switch arr := cur.(type) { + case []any: + if t.index < 0 || t.index >= len(arr) { + return nil, false + } + + return arr[t.index], true + case Array: + return arr.Index(t.index) + default: + return nil, false + } + } + + switch obj := cur.(type) { + case map[string]any: + v, ok := obj[t.field] + + return v, ok + case Object: + return obj.Lookup(t.field) + default: + return nil, false + } +} + +// tokenize splits the part of a path after the leading '$' into tokens. +func tokenize(s string) ([]token, error) { + var toks []token + + for s != "" { + switch s[0] { + case '.': + field, rest := scanField(s[1:]) + if field == "" { + return nil, errorf("empty field name in JSONPath") + } + + toks = append(toks, token{field: field}) + s = rest + case '[': + tok, rest, err := scanBracket(s) + if err != nil { + return nil, err + } + + toks = append(toks, tok) + s = rest + default: + return nil, errorf("unexpected character %q in JSONPath", s[0]) + } + } + + return toks, nil +} + +// scanField reads a dotted field name up to the next '.' or '['. +func scanField(s string) (field, rest string) { + i := strings.IndexAny(s, ".[") + if i < 0 { + return s, "" + } + + return s[:i], s[i:] +} + +// scanBracket reads a "[...]" selector: a numeric index or a quoted field name. +func scanBracket(s string) (token, string, error) { + end := strings.IndexByte(s, ']') + if end < 0 { + return token{}, "", errorf("unterminated '[' in JSONPath") + } + + inner := s[1:end] + rest := s[end+1:] + + if len(inner) >= 2 && (inner[0] == '\'' || inner[0] == '"') { + return token{field: inner[1 : len(inner)-1]}, rest, nil + } + + idx, err := strconv.Atoi(inner) + if err != nil { + return token{}, "", errorf("invalid array index %q in JSONPath", inner) + } + + return token{index: idx, isIndex: true}, rest, nil +} diff --git a/internal/jsonpath/jsonpath_test.go b/internal/jsonpath/jsonpath_test.go new file mode 100644 index 000000000..b86177e60 --- /dev/null +++ b/internal/jsonpath/jsonpath_test.go @@ -0,0 +1,61 @@ +package jsonpath + +import "testing" + +type obj map[string]any + +func (o obj) Lookup(k string) (any, bool) { + v, ok := o[k] + + return v, ok +} + +type arr []any + +func (a arr) Index(i int) (any, bool) { + if i < 0 || i >= len(a) { + return nil, false + } + + return a[i], true +} + +func TestEval(t *testing.T) { + root := map[string]any{"a": map[string]any{"b": []any{1, 2}}, "o": obj{"x": arr{"y"}}} + + cases := []struct { + path string + want any + present bool + }{ + {"$", nil, true}, + {"$.a.b[1]", 2, true}, + {"$['a'].b[0]", 1, true}, + {"$.o.x[0]", "y", true}, + {"$.o.x[5]", nil, false}, + {"$.a.b[9]", nil, false}, + {"$.missing", nil, false}, + {"$.a.b.c", nil, false}, + } + + for _, c := range cases { + got, ok, err := Eval(c.path, root) + if c.path == "$" { + if err != nil || !ok { + t.Errorf("Eval($) = %v %v", ok, err) + } + + continue + } + + if err != nil || ok != c.present || got != c.want { + t.Errorf("Eval(%q) = %v %v %v", c.path, got, ok, err) + } + } + + for _, bad := range []string{"a", "$.a[*]", "$..a", "$.", "$[x]", "$[0", "$x"} { + if _, _, err := Eval(bad, root); err == nil { + t.Errorf("Eval(%q) accepted", bad) + } + } +} diff --git a/internal/vtl/ast.go b/internal/vtl/ast.go new file mode 100644 index 000000000..a927c9e49 --- /dev/null +++ b/internal/vtl/ast.go @@ -0,0 +1,110 @@ +package vtl + +// Template nodes. A template body is a []node rendered in order. +type node any + +// textNode is literal output. +type textNode struct{ text string } + +// refNode prints a reference. quiet is the $!x form. +type refNode struct { + ref *refExpr + quiet bool +} + +// setNode is #set($target = value). +type setNode struct { + target *refExpr + value expr +} + +// ifNode is #if/#elseif/#else/#end. elseBody is nil when there is no #else. +type ifNode struct { + branches []ifBranch + elseBody []node +} + +type ifBranch struct { + cond expr + body []node +} + +// foreachNode is #foreach($varName in iter) body #end. +type foreachNode struct { + varName string + iter expr + body []node +} + +type ( + breakNode struct{} + stopNode struct{} + // returnNode is #return or #return(value). + returnNode struct{ value expr } +) + +// Expressions. +type expr any + +type literal struct{ value any } + +// interpolated is a double-quoted string literal, rendered as a template. +type interpolated struct{ body []node } + +type listExpr struct{ items []expr } + +type rangeExpr struct{ from, to expr } + +type mapExpr struct { + keys []expr + vals []expr +} + +// refExpr is $name followed by a chain of property, method and index steps. +type refExpr struct { + name string + chain []accessor +} + +type accessorKind int + +const ( + accProperty accessorKind = iota + accMethod + accIndex +) + +type accessor struct { + kind accessorKind + name string + args []expr + index expr +} + +type unaryExpr struct { + op string + x expr +} + +type binaryExpr struct { + op string + l, r expr +} + +// Operators, as stored in unaryExpr and binaryExpr. +const ( + opAnd = "&&" + opOr = "||" + opEq = "==" + opNe = "!=" + opLt = "<" + opGt = ">" + opLe = "<=" + opGe = ">=" + opAdd = "+" + opSub = "-" + opMul = "*" + opDiv = "/" + opMod = "%" + opNot = "!" +) diff --git a/internal/vtl/doc.go b/internal/vtl/doc.go new file mode 100644 index 000000000..e5ec08ed6 --- /dev/null +++ b/internal/vtl/doc.go @@ -0,0 +1,11 @@ +// Package vtl is a small Apache Velocity (VTL) engine covering the subset AWS +// mapping templates use: #set, #if/#elseif/#else, #foreach (capped at 1000 +// iterations), #break, #stop, #return, comments, references with property, +// index and method access, literals, ranges, maps and the usual operators, +// plus a Java-like method bridge for strings, lists and maps. +// +// Host values such as API Gateway's $input and $util plug in through Object. +// A null reference renders as an empty string. #macro, #define, #parse, +// #include and #evaluate are rejected at parse time. Every render runs under a +// step budget and the caller's context deadline. +package vtl diff --git a/internal/vtl/edge_test.go b/internal/vtl/edge_test.go new file mode 100644 index 000000000..6e11826f7 --- /dev/null +++ b/internal/vtl/edge_test.go @@ -0,0 +1,129 @@ +package vtl + +import ( + "strings" + "testing" +) + +func TestRenderEdgeCases(t *testing.T) { + cases := []struct { + name, src, want string + }{ + {"unary minus", "#set($a = 3)#set($b = -$a)$b", "-3"}, + {"negative literal", "#set($a = -2.5)$a", "-2.5"}, + {"modulo", "#set($a = 7 % 3)$a", "1"}, + {"float mod", "#set($a = 7.5 % 2)$a", "1.5"}, + {"div zero", "#set($a = 1 / 0)[$a]", "[]"}, + {"float div zero", "#set($a = 1.0 / 0)[$a]", "[]"}, + {"mixed math", "#set($a = 1 + 0.5)#set($b = 2 - 0.5)#set($c = 2 * 0.5)$a $b $c", "1.5 1.5 1.0"}, + {"null arithmetic", "#set($a = $nope + 1)[$a]", "[]"}, + {"string compare", `#if("a" < "b" && "b" >= "b" && "c" > "b" && "a" <= "a")y#end`, "y"}, + {"incomparable", `#if("a" < 1)y#{else}n#end`, "n"}, + {"word comparisons", `#if(1 ne 2 and 2 eq 2.0 and 1 lt 2 and 2 le 2 and 3 ge 2)y#end`, "y"}, + {"or", "#if(false || $nope || true)y#end", "y"}, + {"or word", "#if(false or false)y#{else}n#end", "n"}, + {"null equality", "#if($nope == $other)y#end#if($nope != 1)z#end", "yz"}, + {"bool equality", "#if(true == true && true != false)y#end", "y"}, + {"mixed equality", `#if("1" == 1)y#end`, "y"}, + {"descending range", "#foreach($i in [3..1])$i#end", "321"}, + {"bad range", "#foreach($i in ['a'..2])$i#end[]", "[]"}, + {"foreach non list", "#foreach($i in 5)x#end.", "."}, + {"foreach restores var", "#set($i = 9)#foreach($i in [1..2])#end$i", "9"}, + {"velocityCount", "#foreach($i in [5..6])$velocityCount#end", "12"}, + {"set list index", "#set($l = [1, 2])#set($l[0] = 7)$l", "[7, 2]"}, + {"set map index", `#set($m = {})#set($m["k"] = 1)$m.k`, "1"}, + {"index map and missing", `#set($m = {"a": [1]})$m["a"][0]|$m.a[5]|$m.b.c|$name[0]`, "1|||"}, + { + "string methods", + `#set($s = " Ab ")$s.trim().toLowerCase()|$s.isEmpty()|$s.contains("A")|$s.startsWith(" ")|` + + `$s.endsWith("b")|$s.indexOf("b")|$s.lastIndexOf(" ")|$s.equalsIgnoreCase(" ab ")|` + + `$s.charAt(1)|$s.charAt(99)|$s.substring(1, 3)|$s.substring(9)`, + "ab|false|true|true|false|2|3|true|A||Ab|", + }, + {"replaceFirst", `#set($s = "aXbX")$s.replaceFirst("X", "-")$s.replaceFirst("Z", "-")`, "a-bXaXbX"}, + {"string equals", `#if($name.equals("pet"))y#end`, "y"}, + { + "list methods", + `#set($l = [1, 2, 3])#set($d = $l.remove(0))#set($d = $l.remove(3))$l $l.indexOf(2) $l.isEmpty() ` + + `#set($d = $l.addAll([9]))$l #set($d = $l.set(0, 0))$l [$l.get(9)]`, + "[2, 3] 0 false [2, 3, 9] [0, 3, 9] []", + }, + {"list remove missing", `#set($l = [1])$l.remove("x") [$l.remove(5)]`, "false []"}, + { + "map methods", + `#set($m = {"a": 1})#set($d = $m.putAll({"b": 2}))$m.values() $m.entrySet() $m.remove("a") ` + + `$m.isEmpty() [$m.remove("zz")]`, + "[1, 2] [{key=a, value=1}, {key=b, value=2}] 1 false []", + }, + {"empty property", `#set($l = [])#if($l.empty)y#end`, "y"}, + {"number toString", "$n.toString()", "5"}, + {"unknown method", "[$name.nope()][$n.foo()]", "[][]"}, + {"interpolated stop", `#set($s = "a#stop b")$s`, "a"}, + {"elseif gobble", "#if(false)\na\n#elseif(true)\nb\n#end\n", "b\n"}, + {"return no value", "a#return b", "a"}, + {"block comment gobble", "x\n#* c *#\ny", "x\ny"}, + {"group", "#set($a = (1 + 2) * 3)$a", "9"}, + {"not", "#if(!false && !$nope)y#end", "y"}, + {"null literal", "#set($a = null)[$a]", "[]"}, + {"backslash", `a\b`, `a\b`}, + {"break outside loop", "a#break b", "a"}, + {"hash at end", "a#", "a#"}, + {"dollar at end", "a$", "a$"}, + {"braced unterminated", "${a", "${a"}, + {"hyphenated name", `#set($m = {})#set($m.X-Y = 1)$m`, "{X-Y=1}"}, + } + + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + got := render(t, c.src, map[string]any{"name": "pet", "n": int64(5)}) + if got != c.want { + t.Fatalf("render(%q) = %q, want %q", c.src, got, c.want) + } + }) + } +} + +func TestParseErrors(t *testing.T) { + for _, src := range []string{ + "#* open", "#[[ open", "#{if", "#set(x)", "#set($x 1)", "#set($x = 1", "#if(true)a#{else}b", + "#foreach($a.b in [1])#end", "#foreach($a [1])#end", "#foreach($a in [1]", "#return(1", + "#set($x = [1 2])", "#set($x = [1..2)", "#set($x = {1 2})", "#set($x = {1: 2)", "#set($x = $a.b(1 2))", + "#set($x = $a[1)", "#set($x = @)", "#set($x = foo)", "#set($x = 'open)", "#set($x = 99999999999999999999)", + "#set($x = $)", "#if(true)#elseif(", "#set($x = \"#end\")", "#set($x = (1)", + } { + if _, err := Parse(src); err == nil { + t.Errorf("Parse(%q) succeeded, want error", src) + } else if !strings.Contains(err.Error(), "vtl: parse error") { + t.Errorf("Parse(%q) error = %v", src, err) + } + } +} + +func TestStringifyAndJSON(t *testing.T) { + if Stringify(1e21) != "1000000000000000000000.0" || Stringify(struct{}{}) != "" { + t.Fatalf("Stringify: %s", Stringify(1e21)) + } + + if got := ToJSON(NewList(int64(1), 2.5, true, nil, hostObj{})); got != `[1,2.5,true,null,null]` { + t.Fatalf("ToJSON list = %s", got) + } + + for _, bad := range []string{"", "[1,", `{"a":}`, "]"} { + if _, err := ParseJSON(bad); err == nil { + t.Errorf("ParseJSON(%q) accepted", bad) + } + } + + if v, _ := ParseJSON("1.25"); v != 1.25 { + t.Fatal("float parse") + } + + m := MapOf("a", 1) + if m.Remove("zz") != nil || m.Len() != 1 || len(m.Keys()) != 1 { + t.Fatal("map ops") + } + + if _, ok := NewList().Index(0); ok { + t.Fatal("list index") + } +} diff --git a/internal/vtl/errors.go b/internal/vtl/errors.go new file mode 100644 index 000000000..9ebed0ebc --- /dev/null +++ b/internal/vtl/errors.go @@ -0,0 +1,15 @@ +package vtl + +import "fmt" + +// Error is a template evaluation error, such as an invalid regular expression +// passed to a string method. +type Error struct { + Msg string +} + +func (e *Error) Error() string { return "vtl: " + e.Msg } + +func errorf(format string, args ...any) error { + return &Error{Msg: fmt.Sprintf(format, args...)} +} diff --git a/internal/vtl/eval.go b/internal/vtl/eval.go new file mode 100644 index 000000000..b4a68d454 --- /dev/null +++ b/internal/vtl/eval.go @@ -0,0 +1,724 @@ +package vtl + +import ( + "context" + "errors" + "math" + "strings" +) + +// Execution limits. +const ( + // DefaultMaxSteps bounds the nodes and expressions one render may evaluate. + DefaultMaxSteps = 1_000_000 + // MaxForeachIterations caps every #foreach loop, as AWS does; the loop + // silently stops after this many iterations. + MaxForeachIterations = 1000 + // ctxCheckEvery is how often (in steps) the context deadline is checked. + ctxCheckEvery = 1024 +) + +// ErrStepBudget is returned when a render exceeds its step budget. +var ErrStepBudget = errors.New("vtl: template exceeded its execution step budget") + +// Result is the outcome of a render. +type Result struct { + // Output is the rendered text. + Output string + // Returned reports that the template ran #return. + Returned bool + // ReturnValue is the #return argument, nil when none was given. + ReturnValue any +} + +// RenderOptions tunes a render. A zero MaxSteps selects DefaultMaxSteps. +type RenderOptions struct { + MaxSteps int +} + +// Render evaluates the template with vars as its top-level references. vars is +// modified by #set. Rendering stops with ctx's error when ctx is done. +func (t *Template) Render(ctx context.Context, vars map[string]any, opts RenderOptions) (*Result, error) { + if vars == nil { + vars = map[string]any{} + } + + st := &state{ctx: ctx, vars: vars, maxSteps: opts.MaxSteps} + if st.maxSteps <= 0 { + st.maxSteps = DefaultMaxSteps + } + + err := st.run(t.body) + + var ret *returnSignal + + switch { + case err == nil, errors.Is(err, errStop), errors.Is(err, errBreak): + return &Result{Output: st.out.String()}, nil + case errors.As(err, &ret): + return &Result{Output: st.out.String(), Returned: true, ReturnValue: ret.value}, nil + default: + return nil, err + } +} + +// Control-flow signals, carried as errors through the evaluator. +var ( + errStop = errors.New("vtl: #stop") + errBreak = errors.New("vtl: #break") +) + +type returnSignal struct{ value any } + +func (*returnSignal) Error() string { return "vtl: #return" } + +type state struct { + ctx context.Context + vars map[string]any + out strings.Builder + steps int + maxSteps int +} + +func (s *state) step() error { + s.steps++ + if s.steps > s.maxSteps { + return ErrStepBudget + } + + if s.steps%ctxCheckEvery == 0 { + if err := s.ctx.Err(); err != nil { + return err + } + } + + return nil +} + +func (s *state) run(body []node) error { + for _, n := range body { + if err := s.step(); err != nil { + return err + } + + if err := s.exec(n); err != nil { + return err + } + } + + return nil +} + +func (s *state) exec(n node) error { + switch t := n.(type) { + case *textNode: + s.out.WriteString(t.text) + case *refNode: + v, err := s.evalRef(t.ref) + if err != nil { + return err + } + + s.out.WriteString(Stringify(v)) + case *setNode: + return s.execSet(t) + case *ifNode: + return s.execIf(t) + case *foreachNode: + return s.execForeach(t) + case *breakNode: + return errBreak + case *stopNode: + return errStop + case *returnNode: + return s.execReturn(t) + } + + return nil +} + +func (s *state) execSet(n *setNode) error { + v, err := s.eval(n.value) + if err != nil { + return err + } + + if len(n.target.chain) == 0 { + s.vars[n.target.name] = v + + return nil + } + + parent, err := s.resolveChain(n.target.name, n.target.chain[:len(n.target.chain)-1]) + if err != nil { + return err + } + + last := n.target.chain[len(n.target.chain)-1] + + switch last.kind { + case accProperty: + if m, ok := parent.(*Map); ok { + m.Put(last.name, v) + } + case accIndex: + idx, err := s.eval(last.index) + if err != nil { + return err + } + + setIndex(parent, idx, v) + case accMethod: + return errorf("cannot #set a method call") + } + + return nil +} + +func setIndex(target, idx, v any) { + switch t := target.(type) { + case *Map: + t.Put(Stringify(idx), v) + case *List: + if i, ok := toInt(idx); ok && i >= 0 && i < len(t.Items) { + t.Items[i] = v + } + } +} + +func (s *state) execIf(n *ifNode) error { + for _, b := range n.branches { + v, err := s.eval(b.cond) + if err != nil { + return err + } + + if truthy(v) { + return s.run(b.body) + } + } + + return s.run(n.elseBody) +} + +func (s *state) execForeach(n *foreachNode) error { + src, err := s.eval(n.iter) + if err != nil { + return err + } + + items := iterItems(src) + if len(items) > MaxForeachIterations { + items = items[:MaxForeachIterations] + } + + prevVar, hadVar := s.vars[n.varName] + prevLoop, hadLoop := s.vars["foreach"] + prevCount, hadCount := s.vars["velocityCount"] + + defer func() { + restoreVar(s.vars, n.varName, prevVar, hadVar) + restoreVar(s.vars, "foreach", prevLoop, hadLoop) + restoreVar(s.vars, "velocityCount", prevCount, hadCount) + }() + + for i, it := range items { + s.vars[n.varName] = it + s.vars["foreach"] = MapOf( + "index", int64(i), "count", int64(i+1), "hasNext", i < len(items)-1, + "first", i == 0, "last", i == len(items)-1, + ) + s.vars["velocityCount"] = int64(i + 1) + + err := s.run(n.body) + if errors.Is(err, errBreak) { + return nil + } + + if err != nil { + return err + } + } + + return nil +} + +func restoreVar(vars map[string]any, name string, prev any, had bool) { + if had { + vars[name] = prev + } else { + delete(vars, name) + } +} + +// iterItems returns what #foreach walks: a list's items, a map's values, or +// nothing. +func iterItems(v any) []any { + switch t := v.(type) { + case *List: + return append([]any(nil), t.Items...) + case *Map: + out := make([]any, 0, t.Len()) + for _, k := range t.keys { + out = append(out, t.vals[k]) + } + + return out + default: + return nil + } +} + +func (s *state) execReturn(n *returnNode) error { + if n.value == nil { + return &returnSignal{} + } + + v, err := s.eval(n.value) + if err != nil { + return err + } + + return &returnSignal{value: v} +} + +func (s *state) eval(e expr) (any, error) { + if err := s.step(); err != nil { + return nil, err + } + + switch t := e.(type) { + case *literal: + return t.value, nil + case *interpolated: + return s.evalInterpolated(t) + case *refExpr: + return s.evalRef(t) + case *unaryExpr: + return s.evalUnary(t) + case *binaryExpr: + return s.evalBinary(t) + } + + return s.evalCollection(e) +} + +func (s *state) evalCollection(e expr) (any, error) { + switch t := e.(type) { + case *listExpr: + return s.evalList(t) + case *rangeExpr: + return s.evalRange(t) + case *mapExpr: + return s.evalMap(t) + } + + return nil, errorf("unknown expression %T", e) +} + +// evalInterpolated renders a double-quoted string as a template sharing the +// caller's variables and step budget. +func (s *state) evalInterpolated(t *interpolated) (any, error) { + sub := &state{ctx: s.ctx, vars: s.vars, steps: s.steps, maxSteps: s.maxSteps} + err := sub.run(t.body) + s.steps = sub.steps + + if err != nil && !errors.Is(err, errStop) { + return nil, err + } + + return sub.out.String(), nil +} + +func (s *state) evalList(t *listExpr) (any, error) { + l := NewList() + + for _, it := range t.items { + v, err := s.eval(it) + if err != nil { + return nil, err + } + + l.Items = append(l.Items, v) + } + + return l, nil +} + +func (s *state) evalRange(t *rangeExpr) (any, error) { + from, err := s.eval(t.from) + if err != nil { + return nil, err + } + + to, err := s.eval(t.to) + if err != nil { + return nil, err + } + + a, okA := toInt(from) + b, okB := toInt(to) + + if !okA || !okB { + return nil, nil + } + + l := NewList() + + step := 1 + if b < a { + step = -1 + } + + for i := a; ; i += step { + l.Items = append(l.Items, int64(i)) + + if i == b || len(l.Items) > MaxForeachIterations { + break + } + } + + return l, nil +} + +func (s *state) evalMap(t *mapExpr) (any, error) { + m := NewMap() + + for i := range t.keys { + k, err := s.eval(t.keys[i]) + if err != nil { + return nil, err + } + + v, err := s.eval(t.vals[i]) + if err != nil { + return nil, err + } + + m.Put(Stringify(k), v) + } + + return m, nil +} + +func (s *state) evalUnary(t *unaryExpr) (any, error) { + v, err := s.eval(t.x) + if err != nil { + return nil, err + } + + if t.op == opNot { + return !truthy(v), nil + } + + switch n := v.(type) { + case int64: + return -n, nil + case float64: + return -n, nil + default: + return nil, nil + } +} + +func (s *state) evalBinary(t *binaryExpr) (any, error) { + l, err := s.eval(t.l) + if err != nil { + return nil, err + } + + // && and || short-circuit. + if t.op == opAnd || t.op == opOr { + if truthy(l) == (t.op == opOr) { + return t.op == opOr, nil + } + + r, rerr := s.eval(t.r) + + return truthy(r), rerr + } + + r, err := s.eval(t.r) + if err != nil { + return nil, err + } + + switch t.op { + case opEq: + return equal(l, r), nil + case opNe: + return !equal(l, r), nil + case opLt, opGt, opLe, opGe: + return compare(t.op, l, r), nil + default: + return arith(t.op, l, r), nil + } +} + +// evalRef resolves a reference and its chain. A missing value is nil. +func (s *state) evalRef(r *refExpr) (any, error) { + return s.resolveChain(r.name, r.chain) +} + +func (s *state) resolveChain(name string, chain []accessor) (any, error) { + cur, ok := s.vars[name] + if !ok { + return nil, nil + } + + for _, a := range chain { + if cur == nil { + return nil, nil + } + + var err error + if cur, err = s.access(cur, a); err != nil { + return nil, err + } + } + + return cur, nil +} + +// access applies one accessor step to cur. +func (s *state) access(cur any, a accessor) (any, error) { + switch a.kind { + case accProperty: + return property(cur, a.name), nil + case accIndex: + idx, err := s.eval(a.index) + if err != nil { + return nil, err + } + + return index(cur, idx), nil + case accMethod: + args := make([]any, len(a.args)) + + for i, ae := range a.args { + v, err := s.eval(ae) + if err != nil { + return nil, err + } + + args[i] = v + } + + return callMethod(cur, a.name, args) + } + + return nil, nil +} + +func property(v any, name string) any { + switch t := v.(type) { + case *Map: + got, _ := t.Get(name) + + return got + case Object: + got, _ := t.Get(name) + + return got + } + + // Bean-style getters: $list.empty, $str.empty. + if name == "empty" { + if r, err := callMethod(v, mIsEmpty, nil); err == nil { + return r + } + } + + return nil +} + +func index(v, idx any) any { + switch t := v.(type) { + case *List: + i, ok := toInt(idx) + if !ok || i < 0 || i >= len(t.Items) { + return nil + } + + return t.Items[i] + case *Map: + got, _ := t.Get(Stringify(idx)) + + return got + case Object: + got, _ := t.Get(Stringify(idx)) + + return got + } + + return nil +} + +// truthy is Velocity's #if test: null and false are false, anything else true. +func truthy(v any) bool { + switch t := v.(type) { + case nil: + return false + case bool: + return t + default: + return true + } +} + +func toInt(v any) (int, bool) { + switch t := v.(type) { + case int64: + return int(t), true + case float64: + return int(t), true + default: + return 0, false + } +} + +func toFloat(v any) (float64, bool) { + switch t := v.(type) { + case int64: + return float64(t), true + case float64: + return t, true + default: + return 0, false + } +} + +func equal(l, r any) bool { + if l == nil || r == nil { + return l == nil && r == nil + } + + if a, ok := toFloat(l); ok { + if b, ok := toFloat(r); ok { + return a == b + } + } + + if a, ok := l.(bool); ok { + b, ok := r.(bool) + + return ok && a == b + } + + // Different types compare by their string form, as Velocity does. + return Stringify(l) == Stringify(r) +} + +func compare(op string, l, r any) bool { + var c int + + a, okA := toFloat(l) + b, okB := toFloat(r) + + switch { + case okA && okB: + c = cmpFloat(a, b) + default: + ls, isLS := l.(string) + rs, isRS := r.(string) + + if !isLS || !isRS { + return false + } + + c = strings.Compare(ls, rs) + } + + switch op { + case opLt: + return c < 0 + case opGt: + return c > 0 + case opLe: + return c <= 0 + default: + return c >= 0 + } +} + +func cmpFloat(a, b float64) int { + switch { + case a < b: + return -1 + case a > b: + return 1 + default: + return 0 + } +} + +// arith applies + - * / %. A string operand of + concatenates; integer +// operands keep integer arithmetic; a division by zero is null. +func arith(op string, l, r any) any { + if op == opAdd { + _, ls := l.(string) + _, rs := r.(string) + + if ls || rs { + return Stringify(l) + Stringify(r) + } + } + + li, lInt := l.(int64) + ri, rInt := r.(int64) + + if lInt && rInt { + return intArith(op, li, ri) + } + + a, okA := toFloat(l) + b, okB := toFloat(r) + + if !okA || !okB { + return nil + } + + return floatArith(op, a, b) +} + +func floatArith(op string, a, b float64) any { + switch op { + case opAdd: + return a + b + case opSub: + return a - b + case opMul: + return a * b + } + + if b == 0 { + return nil + } + + if op == opDiv { + return a / b + } + + return math.Mod(a, b) +} + +func intArith(op string, a, b int64) any { + switch op { + case opAdd: + return a + b + case opSub: + return a - b + case opMul: + return a * b + } + + if b == 0 { + return nil + } + + if op == opDiv { + return a / b + } + + return a % b +} diff --git a/internal/vtl/methods.go b/internal/vtl/methods.go new file mode 100644 index 000000000..6475baabb --- /dev/null +++ b/internal/vtl/methods.go @@ -0,0 +1,430 @@ +package vtl + +import ( + "regexp" + "strings" +) + +// Method names shared by more than one receiver type. +const ( + mIsEmpty = "isEmpty" + mContains = "contains" + mIndexOf = "indexOf" + mSize = "size" + mGet = "get" + mRemove = "remove" + mSet = "set" + mPut = "put" +) + +// pairArgs is the argument count of a two-argument method such as put or set. +const pairArgs = 2 + +type ( + stringFn func(s string, args []any) (any, error) + listFn func(l *List, args []any) any + mapFn func(m *Map, args []any) any +) + +// The Java-like method bridge for strings, lists and maps. +// +//nolint:gochecknoglobals // read-only dispatch tables +var ( + stringMethods = map[string]stringFn{ + "length": func(s string, _ []any) (any, error) { return int64(len([]rune(s))), nil }, + mIsEmpty: func(s string, _ []any) (any, error) { return s == "", nil }, + mContains: strPredicate(strings.Contains), + "startsWith": strPredicate(strings.HasPrefix), + "endsWith": strPredicate(strings.HasSuffix), + "equalsIgnoreCase": strPredicate(strings.EqualFold), + mIndexOf: strIndex(strings.Index), + "lastIndexOf": strIndex(strings.LastIndex), + "substring": func(s string, args []any) (any, error) { return substring(s, args), nil }, + "replace": strReplace, + "replaceAll": func(s string, args []any) (any, error) { return regexReplace(s, true, args) }, + "replaceFirst": func(s string, args []any) (any, error) { return regexReplace(s, false, args) }, + "split": strSplit, + "toLowerCase": func(s string, _ []any) (any, error) { return strings.ToLower(s), nil }, + "toUpperCase": func(s string, _ []any) (any, error) { return strings.ToUpper(s), nil }, + "trim": func(s string, _ []any) (any, error) { return strings.TrimSpace(s), nil }, + "matches": strMatches, + "charAt": strCharAt, + } + + listMethods = map[string]listFn{ + mSize: func(l *List, _ []any) any { return int64(len(l.Items)) }, + mIsEmpty: func(l *List, _ []any) any { return len(l.Items) == 0 }, + mGet: listGet, + "add": listAdd, + "addAll": listAddAll, + mContains: func(l *List, args []any) any { return len(args) == 1 && listIndexOf(l, args[0]) >= 0 }, + mIndexOf: listIndexOfMethod, + mRemove: listRemove, + mSet: listSet, + } + + mapMethods = map[string]mapFn{ + mGet: mapGet, + mPut: mapPut, + "putAll": mapPutAll, + "containsKey": mapContainsKey, + mRemove: func(m *Map, args []any) any { return m.Remove(Stringify(firstArg(args))) }, + "keySet": func(m *Map, _ []any) any { return stringList(m.keys) }, + "values": func(m *Map, _ []any) any { return NewList(iterItems(m)...) }, + "entrySet": mapEntrySet, + mSize: func(m *Map, _ []any) any { return int64(m.Len()) }, + mIsEmpty: func(m *Map, _ []any) any { return m.Len() == 0 }, + } +) + +// callMethod dispatches a method call to the bridge for strings, lists and +// maps, or to a host Object. An unknown method yields null, as Velocity does +// when no method matches. +func callMethod(v any, name string, args []any) (any, error) { + switch { + case name == "toString" && len(args) == 0: + return Stringify(v), nil + case name == "equals" && len(args) == 1: + return equal(v, args[0]), nil + } + + switch t := v.(type) { + case string: + if fn, ok := stringMethods[name]; ok { + return fn(t, args) + } + case Object: + return callObject(t, name, args) + default: + return callCollection(v, name, args), nil + } + + return nil, nil +} + +func callCollection(v any, name string, args []any) any { + switch t := v.(type) { + case *List: + if fn, ok := listMethods[name]; ok { + return fn(t, args) + } + case *Map: + if fn, ok := mapMethods[name]; ok { + return fn(t, args) + } + } + + return nil +} + +func callObject(o Object, name string, args []any) (any, error) { + r, ok, err := o.Call(name, args) + if err != nil || !ok { + return nil, err + } + + return r, nil +} + +// strArg returns argument i as a string; a non-string is stringified. +func strArg(args []any, i int) (string, bool) { + if i >= len(args) || args[i] == nil { + return "", false + } + + return Stringify(args[i]), true +} + +func intArg(args []any, i int) (int, bool) { + if i >= len(args) { + return 0, false + } + + return toInt(args[i]) +} + +func firstArg(args []any) any { + if len(args) == 0 { + return nil + } + + return args[0] +} + +func strPredicate(pred func(s, arg string) bool) stringFn { + return func(s string, args []any) (any, error) { + a, ok := strArg(args, 0) + + return ok && pred(s, a), nil + } +} + +// strIndex converts a byte offset to a character offset (-1 stays -1). +func strIndex(find func(s, sub string) int) stringFn { + return func(s string, args []any) (any, error) { + a, _ := strArg(args, 0) + + i := find(s, a) + if i < 0 { + return int64(-1), nil + } + + return int64(len([]rune(s[:i]))), nil + } +} + +func strReplace(s string, args []any) (any, error) { + from, ok1 := strArg(args, 0) + to, ok2 := strArg(args, 1) + + if !ok1 || !ok2 { + return nil, nil + } + + return strings.ReplaceAll(s, from, to), nil +} + +func strMatches(s string, args []any) (any, error) { + pattern, _ := strArg(args, 0) + + re, err := compile("^(?:" + pattern + ")$") + if err != nil { + return nil, err + } + + return re.MatchString(s), nil +} + +func strCharAt(s string, args []any) (any, error) { + i, ok := intArg(args, 0) + r := []rune(s) + + if !ok || i < 0 || i >= len(r) { + return nil, nil + } + + return string(r[i]), nil +} + +func substring(s string, args []any) any { + r := []rune(s) + + begin, ok := intArg(args, 0) + if !ok || begin < 0 || begin > len(r) { + return nil + } + + end := len(r) + + if e, ok := intArg(args, 1); ok { + if e < begin || e > len(r) { + return nil + } + + end = e + } + + return string(r[begin:end]) +} + +func compile(pattern string) (*regexp.Regexp, error) { + re, err := regexp.Compile(pattern) + if err != nil { + return nil, errorf("invalid regular expression %q: %v", pattern, err) + } + + return re, nil +} + +func regexReplace(s string, all bool, args []any) (any, error) { + pattern, _ := strArg(args, 0) + repl, _ := strArg(args, 1) + + re, err := compile(pattern) + if err != nil { + return nil, err + } + + if all { + return re.ReplaceAllString(s, repl), nil + } + + loc := re.FindStringSubmatchIndex(s) + if loc == nil { + return s, nil + } + + return s[:loc[0]] + string(re.ExpandString(nil, repl, s, loc)) + s[loc[1]:], nil +} + +// strSplit follows Java's String.split: the argument is a regex and trailing +// empty strings are dropped. +func strSplit(s string, args []any) (any, error) { + pattern, _ := strArg(args, 0) + + re, err := compile(pattern) + if err != nil { + return nil, err + } + + parts := re.Split(s, -1) + for len(parts) > 0 && parts[len(parts)-1] == "" { + parts = parts[:len(parts)-1] + } + + return stringList(parts), nil +} + +func listGet(l *List, args []any) any { + v, _ := l.Index(intOr(args, -1)) + + return v +} + +func intOr(args []any, def int) int { + if i, ok := intArg(args, 0); ok { + return i + } + + return def +} + +func listAdd(l *List, args []any) any { + if len(args) != 1 { + return nil + } + + l.Items = append(l.Items, args[0]) + + return true +} + +func listAddAll(l *List, args []any) any { + other, ok := firstArg(args).(*List) + if !ok { + return false + } + + l.Items = append(l.Items, other.Items...) + + return true +} + +func listIndexOfMethod(l *List, args []any) any { + if len(args) != 1 { + return nil + } + + return int64(listIndexOf(l, args[0])) +} + +func listIndexOf(l *List, v any) int { + for i, it := range l.Items { + if equal(it, v) { + return i + } + } + + return -1 +} + +// listRemove follows Java: remove(int) drops by index, remove(Object) by value. +func listRemove(l *List, args []any) any { + if len(args) != 1 { + return nil + } + + if i, ok := args[0].(int64); ok { + if i < 0 || int(i) >= len(l.Items) { + return nil + } + + prev := l.Items[i] + l.Items = append(l.Items[:i], l.Items[i+1:]...) + + return prev + } + + idx := listIndexOf(l, args[0]) + if idx < 0 { + return false + } + + l.Items = append(l.Items[:idx], l.Items[idx+1:]...) + + return true +} + +func listSet(l *List, args []any) any { + i, ok := intArg(args, 0) + if !ok || len(args) != pairArgs || i < 0 || i >= len(l.Items) { + return nil + } + + prev := l.Items[i] + l.Items[i] = args[1] + + return prev +} + +func mapGet(m *Map, args []any) any { + key, ok := strArg(args, 0) + if !ok { + return nil + } + + v, _ := m.Get(key) + + return v +} + +func mapPut(m *Map, args []any) any { + if len(args) != pairArgs { + return nil + } + + return m.Put(Stringify(args[0]), args[1]) +} + +func mapPutAll(m *Map, args []any) any { + if other, ok := firstArg(args).(*Map); ok { + for _, k := range other.keys { + m.Put(k, other.vals[k]) + } + } + + return nil +} + +func mapContainsKey(m *Map, args []any) any { + key, ok := strArg(args, 0) + if !ok { + return false + } + + _, found := m.Get(key) + + return found +} + +func mapEntrySet(m *Map, _ []any) any { + l := NewList() + + for _, k := range m.keys { + entry := NewMap() + entry.Put("key", k) + entry.Put("value", m.vals[k]) + l.Items = append(l.Items, entry) + } + + return l +} + +func stringList(ss []string) *List { + l := NewList() + for _, s := range ss { + l.Items = append(l.Items, s) + } + + return l +} diff --git a/internal/vtl/parse.go b/internal/vtl/parse.go new file mode 100644 index 000000000..c147cf6c7 --- /dev/null +++ b/internal/vtl/parse.go @@ -0,0 +1,1109 @@ +package vtl + +import ( + "fmt" + "strconv" + "strings" +) + +// Template is a parsed VTL template. +type Template struct { + body []node +} + +// ParseError reports a template that could not be parsed. +type ParseError struct { + Pos int + Msg string +} + +func (e *ParseError) Error() string { + return fmt.Sprintf("vtl: parse error at offset %d: %s", e.Pos, e.Msg) +} + +// Parse parses src into a Template. +func Parse(src string) (*Template, error) { + p := &parser{src: src} + + body, term, err := p.parseBlock() + if err != nil { + return nil, err + } + + if term != "" { + return nil, p.errf("unexpected #%s", term) + } + + return &Template{body: body}, nil +} + +type parser struct { + src string + pos int +} + +func (p *parser) errf(format string, args ...any) error { + return &ParseError{Pos: p.pos, Msg: fmt.Sprintf(format, args...)} +} + +// Directive names. +const ( + dirSet = "set" + dirIf = "if" + dirElseIf = "elseif" + dirElse = "else" + dirEnd = "end" + dirForeach = "foreach" + dirBreak = "break" + dirStop = "stop" + dirReturn = "return" +) + +// knownDirectives are the directives the subset implements. +var knownDirectives = map[string]bool{ //nolint:gochecknoglobals // read-only lookup table + dirSet: true, dirIf: true, dirElseIf: true, dirElse: true, dirEnd: true, + dirForeach: true, dirBreak: true, dirStop: true, dirReturn: true, +} + +// Directives the subset rejects outright. +var unsupportedDirectives = map[string]bool{ //nolint:gochecknoglobals // read-only lookup table + "macro": true, "define": true, "parse": true, "include": true, "evaluate": true, +} + +// parseBlock parses nodes until EOF or a block terminator (#else, #elseif, +// #end), which it returns. For #elseif the parser is left at its condition. +func (p *parser) parseBlock() ([]node, string, error) { + var nodes []node + + for p.pos < len(p.src) { + i := strings.IndexAny(p.src[p.pos:], "$#\\") + if i < 0 { + nodes = appendText(nodes, p.src[p.pos:]) + p.pos = len(p.src) + + break + } + + nodes = appendText(nodes, p.src[p.pos:p.pos+i]) + p.pos += i + + switch p.src[p.pos] { + case '\\': + nodes = p.parseEscape(nodes) + case '$': + var err error + + nodes, err = p.parseOutputRef(nodes) + if err != nil { + return nil, "", err + } + default: + var ( + term string + err error + ) + + nodes, term, err = p.parseHash(nodes) + if err != nil || term != "" { + return nodes, term, err + } + } + } + + return nodes, "", nil +} + +func appendText(nodes []node, s string) []node { + if s == "" { + return nodes + } + + if n := len(nodes); n > 0 { + if t, ok := nodes[n-1].(*textNode); ok { + t.text += s + + return nodes + } + } + + return append(nodes, &textNode{text: s}) +} + +// parseEscape handles a backslash: \$ and \# print the next character +// literally; any other backslash is plain text. +func (p *parser) parseEscape(nodes []node) []node { + if p.pos+1 < len(p.src) && (p.src[p.pos+1] == '$' || p.src[p.pos+1] == '#') { + nodes = appendText(nodes, p.src[p.pos+1:p.pos+2]) + p.pos += 2 + + return nodes + } + + p.pos++ + + return appendText(nodes, `\`) +} + +// parseOutputRef parses a reference in text. A '$' that does not start a +// reference is plain text. +func (p *parser) parseOutputRef(nodes []node) ([]node, error) { + start := p.pos + + ref, quiet, ok, err := p.parseRef() + if err != nil { + return nil, err + } + + if !ok { + p.pos = start + 1 + + return appendText(nodes, "$"), nil + } + + return append(nodes, &refNode{ref: ref, quiet: quiet}), nil +} + +// parseHash handles '#': comments, unparsed blocks and directives. A '#' that +// starts none of them is plain text. +func (p *parser) parseHash(nodes []node) ([]node, string, error) { + start := p.pos + rest := p.src[p.pos:] + + switch { + case strings.HasPrefix(rest, "##"): + end := strings.IndexByte(rest, '\n') + if end < 0 { + p.pos = len(p.src) + } else { + p.pos += end + 1 + } + + return nodes, "", nil + case strings.HasPrefix(rest, "#*"): + end := strings.Index(rest[2:], "*#") + if end < 0 { + return nil, "", p.errf("unterminated #* comment") + } + + p.pos += end + 4 + + return p.gobble(nodes, start), "", nil + case strings.HasPrefix(rest, "#[["): + end := strings.Index(rest, "]]#") + if end < 0 { + return nil, "", p.errf("unterminated #[[ block") + } + + p.pos += end + 3 + + return appendText(nodes, rest[3:end]), "", nil + } + + name, braced := p.directiveName() + if unsupportedDirectives[name] { + return nil, "", p.errf("#%s is not supported", name) + } + + return p.parseDirective(nodes, start, name, braced) +} + +// directiveName reads the identifier after '#' (optionally #{name}) without +// consuming it. +func (p *parser) directiveName() (string, bool) { + i := p.pos + 1 + braced := i < len(p.src) && p.src[i] == '{' + + if braced { + i++ + } + + j := i + for j < len(p.src) && isIdentChar(p.src[j]) { + j++ + } + + return p.src[i:j], braced +} + +// consumeDirectiveName advances past #name or #{name}. +func (p *parser) consumeDirectiveName(name string, braced bool) bool { + p.pos++ + + if braced { + p.pos++ + } + + p.pos += len(name) + + if braced { + if p.pos >= len(p.src) || p.src[p.pos] != '}' { + return false + } + + p.pos++ + } + + return true +} + +func (p *parser) parseDirective(nodes []node, start int, name string, braced bool) ([]node, string, error) { + if !knownDirectives[name] { + p.pos++ + + return appendText(nodes, "#"), "", nil + } + + if !p.consumeDirectiveName(name, braced) { + return nil, "", p.errf("malformed #{%s}", name) + } + + switch name { + case dirElse, dirEnd: + return p.gobble(nodes, start), name, nil + case dirElseIf: + // The caller parses the condition; gobbling waits until it has. + return nodes, name, nil + case dirIf: + return p.parseIf(nodes, start) + case dirForeach: + return p.parseForeach(nodes, start) + case dirReturn: + return p.parseReturn(nodes, start) + } + + n, err := p.parseSimpleDirective(name) + if err != nil { + return nil, "", err + } + + return append(p.gobble(nodes, start), n), "", nil +} + +// parseSimpleDirective parses the directives that produce a single node. +func (p *parser) parseSimpleDirective(name string) (node, error) { + switch name { + case dirBreak: + return &breakNode{}, nil + case dirStop: + return &stopNode{}, nil + default: + return p.parseSet() + } +} + +// gobble drops the whitespace-only line around a directive that sits alone on +// its line: the indentation before it (already emitted as text) and the +// trailing whitespace and newline after it. +func (p *parser) gobble(nodes []node, start int) []node { + end, ok := p.blankLineEnd(start) + if !ok { + return nodes + } + + p.pos = end + + if n := len(nodes); n > 0 { + if t, ok := nodes[n-1].(*textNode); ok { + t.text = strings.TrimRight(t.text, " \t") + } + } + + return nodes +} + +// blankLineEnd reports whether the directive spanning start..p.pos is alone on +// its line, and returns the offset just past that line's newline. +func (p *parser) blankLineEnd(start int) (int, bool) { + lineStart := strings.LastIndexByte(p.src[:start], '\n') + 1 + if strings.TrimLeft(p.src[lineStart:start], " \t") != "" { + return 0, false + } + + rest := p.src[p.pos:] + nl := strings.IndexByte(rest, '\n') + + if nl < 0 { + return len(p.src), strings.TrimLeft(rest, " \t\r") == "" + } + + return p.pos + nl + 1, strings.TrimLeft(rest[:nl], " \t\r") == "" +} + +func (p *parser) parseSet() (node, error) { + if err := p.expectOpenParen(); err != nil { + return nil, err + } + + p.skipSpace() + + target, _, ok, err := p.parseRef() + if err != nil { + return nil, err + } + + if !ok { + return nil, p.errf("#set needs a reference") + } + + p.skipSpace() + + if !p.consume("=") { + return nil, p.errf("#set needs '='") + } + + value, err := p.parseExpr() + if err != nil { + return nil, err + } + + if err := p.expectCloseParen(); err != nil { + return nil, err + } + + return &setNode{target: target, value: value}, nil +} + +func (p *parser) parseCondition() (expr, error) { + if err := p.expectOpenParen(); err != nil { + return nil, err + } + + cond, err := p.parseExpr() + if err != nil { + return nil, err + } + + if err := p.expectCloseParen(); err != nil { + return nil, err + } + + return cond, nil +} + +func (p *parser) parseIf(nodes []node, start int) ([]node, string, error) { + cond, err := p.parseCondition() + if err != nil { + return nil, "", err + } + + nodes = p.gobble(nodes, start) + n := &ifNode{} + + for { + body, term, berr := p.parseBlock() + if berr != nil { + return nil, "", berr + } + + n.branches = append(n.branches, ifBranch{cond: cond, body: body}) + + switch term { + case dirElseIf: + elseStart := strings.LastIndex(p.src[:p.pos], "#") + + if cond, err = p.parseCondition(); err != nil { + return nil, "", err + } + + p.gobbleTrailing(&n.branches[len(n.branches)-1].body, elseStart) + + continue + case dirElse: + if n.elseBody, err = p.parseElse(); err != nil { + return nil, "", err + } + case dirEnd: + default: + return nil, "", p.errf("#if without #end") + } + + return append(nodes, n), "", nil + } +} + +// parseElse parses an #else body up to its #end. The result is never nil, so +// the evaluator can tell an empty #else from none. +func (p *parser) parseElse() ([]node, error) { + body, term, err := p.parseBlock() + if err != nil { + return nil, err + } + + if term != dirEnd { + return nil, p.errf("#else without #end") + } + + if body == nil { + body = []node{} + } + + return body, nil +} + +// gobbleTrailing applies gobble to a body that has already been closed. +func (p *parser) gobbleTrailing(body *[]node, start int) { + *body = p.gobble(*body, start) +} + +func (p *parser) parseForeach(nodes []node, start int) ([]node, string, error) { + if err := p.expectOpenParen(); err != nil { + return nil, "", err + } + + p.skipSpace() + + ref, _, ok, err := p.parseRef() + if err != nil { + return nil, "", err + } + + if !ok || len(ref.chain) > 0 { + return nil, "", p.errf("#foreach needs a plain loop variable") + } + + p.skipSpace() + + if !p.consumeWord("in") { + return nil, "", p.errf("#foreach needs 'in'") + } + + iter, err := p.parseExpr() + if err != nil { + return nil, "", err + } + + if cerr := p.expectCloseParen(); cerr != nil { + return nil, "", cerr + } + + nodes = p.gobble(nodes, start) + + body, term, err := p.parseBlock() + if err != nil { + return nil, "", err + } + + if term != dirEnd { + return nil, "", p.errf("#foreach without #end") + } + + return append(nodes, &foreachNode{varName: ref.name, iter: iter, body: body}), "", nil +} + +func (p *parser) parseReturn(nodes []node, start int) ([]node, string, error) { + n := &returnNode{} + + save := p.pos + p.skipSpace() + + if p.peek() != '(' { + p.pos = save + + return append(p.gobble(nodes, start), n), "", nil + } + + p.pos++ + p.skipSpace() + + if p.peek() != ')' { + v, err := p.parseExpr() + if err != nil { + return nil, "", err + } + + n.value = v + } + + if err := p.expectCloseParen(); err != nil { + return nil, "", err + } + + return append(p.gobble(nodes, start), n), "", nil +} + +// parseRef parses $name, $!name, ${name} and $!{name} with their accessor +// chain. ok is false (and the position unchanged) when the text at the cursor +// is not a reference. +func (p *parser) parseRef() (ref *refExpr, quiet, ok bool, err error) { + start := p.pos + i := p.pos + 1 + + if i < len(p.src) && p.src[i] == '!' { + quiet = true + i++ + } + + braced := i < len(p.src) && p.src[i] == '{' + if braced { + i++ + } + + if i >= len(p.src) || !isIdentStart(p.src[i]) { + return nil, false, false, nil + } + + p.pos = i + ref = &refExpr{name: p.ident()} + + if err := p.parseChain(ref); err != nil { + return nil, false, false, err + } + + if braced { + if p.peek() != '}' { + p.pos = start + + return nil, false, false, nil + } + + p.pos++ + } + + return ref, quiet, true, nil +} + +// parseChain reads .prop, .method(args) and [index] steps. +func (p *parser) parseChain(ref *refExpr) error { + for p.pos < len(p.src) { + switch c := p.src[p.pos]; { + case c == '.' && p.pos+1 < len(p.src) && isIdentStart(p.src[p.pos+1]): + p.pos++ + name := p.ident() + + if p.peek() != '(' { + ref.chain = append(ref.chain, accessor{kind: accProperty, name: name}) + + continue + } + + p.pos++ + + args, err := p.parseArgs(')') + if err != nil { + return err + } + + ref.chain = append(ref.chain, accessor{kind: accMethod, name: name, args: args}) + case c == '[': + p.pos++ + + idx, err := p.parseExpr() + if err != nil { + return err + } + + p.skipSpace() + + if !p.consume("]") { + return p.errf("expected ']'") + } + + ref.chain = append(ref.chain, accessor{kind: accIndex, index: idx}) + default: + return nil + } + } + + return nil +} + +// parseArgs parses a comma-separated expression list up to closer. +func (p *parser) parseArgs(closer byte) ([]expr, error) { + var args []expr + + p.skipSpace() + + if p.peek() == closer { + p.pos++ + + return args, nil + } + + for { + a, err := p.parseExpr() + if err != nil { + return nil, err + } + + args = append(args, a) + + p.skipSpace() + + switch p.peek() { + case ',': + p.pos++ + case closer: + p.pos++ + + return args, nil + default: + return nil, p.errf("expected ',' or '%c'", closer) + } + } +} + +// Expression parsing, lowest precedence first. + +func (p *parser) parseExpr() (expr, error) { return p.parseOr() } + +func (p *parser) parseOr() (expr, error) { + return p.parseBinary(p.parseAnd, map[string]string{opOr: opOr, "or": opOr}) +} + +func (p *parser) parseAnd() (expr, error) { + return p.parseBinary(p.parseEquality, map[string]string{opAnd: opAnd, "and": opAnd}) +} + +func (p *parser) parseEquality() (expr, error) { + return p.parseBinary(p.parseRelational, map[string]string{opEq: opEq, opNe: opNe, "eq": opEq, "ne": opNe}) +} + +func (p *parser) parseRelational() (expr, error) { + return p.parseBinary(p.parseAdditive, map[string]string{ + opLe: opLe, opGe: opGe, opLt: opLt, opGt: opGt, "le": opLe, "ge": opGe, "lt": opLt, "gt": opGt, + }) +} + +func (p *parser) parseAdditive() (expr, error) { + return p.parseBinary(p.parseMultiplicative, map[string]string{opAdd: opAdd, opSub: opSub}) +} + +func (p *parser) parseMultiplicative() (expr, error) { + return p.parseBinary(p.parseUnary, map[string]string{opMul: opMul, opDiv: opDiv, opMod: opMod}) +} + +// parseBinary parses a left-associative chain of the operators in ops. +func (p *parser) parseBinary(next func() (expr, error), ops map[string]string) (expr, error) { + l, err := next() + if err != nil { + return nil, err + } + + for { + op, ok := p.matchOperator(ops) + if !ok { + return l, nil + } + + r, err := next() + if err != nil { + return nil, err + } + + l = &binaryExpr{op: op, l: l, r: r} + } +} + +// matchOperator consumes the longest operator in ops at the cursor. Word +// operators must not be followed by an identifier character. +func (p *parser) matchOperator(ops map[string]string) (string, bool) { + p.skipSpace() + + best := "" + + for tok := range ops { + if !strings.HasPrefix(p.src[p.pos:], tok) || len(tok) <= len(best) { + continue + } + + end := p.pos + len(tok) + if isIdentStart(tok[0]) && end < len(p.src) && isIdentChar(p.src[end]) { + continue + } + + best = tok + } + + if best == "" { + return "", false + } + + p.pos += len(best) + + return ops[best], true +} + +func (p *parser) parseUnary() (expr, error) { + p.skipSpace() + + switch { + case p.peek() == '!' && !strings.HasPrefix(p.src[p.pos:], opNe): + p.pos++ + + x, err := p.parseUnary() + if err != nil { + return nil, err + } + + return &unaryExpr{op: opNot, x: x}, nil + case p.consumeWord("not"): + x, err := p.parseUnary() + if err != nil { + return nil, err + } + + return &unaryExpr{op: opNot, x: x}, nil + case p.peek() == '-' && p.pos+1 < len(p.src) && !isDigit(p.src[p.pos+1]): + p.pos++ + + x, err := p.parseUnary() + if err != nil { + return nil, err + } + + return &unaryExpr{op: opSub, x: x}, nil + } + + return p.parsePrimary() +} + +func (p *parser) parsePrimary() (expr, error) { + p.skipSpace() + + if p.pos >= len(p.src) { + return nil, p.errf("unexpected end of template in expression") + } + + switch c := p.src[p.pos]; { + case c == '$': + ref, _, ok, err := p.parseRef() + if err != nil { + return nil, err + } + + if !ok { + return nil, p.errf("invalid reference") + } + + return ref, nil + case c == '"': + return p.parseDoubleQuoted() + case c == '\'': + return p.parseSingleQuoted() + case isDigit(c) || c == '-': + return p.parseNumber() + case isIdentStart(c): + return p.parseKeyword() + default: + return p.parseBracketed(c) + } +} + +// parseBracketed parses a list or range, a map literal or a parenthesised +// expression. +func (p *parser) parseBracketed(c byte) (expr, error) { + switch c { + case '[': + p.pos++ + + return p.parseListOrRange() + case '{': + p.pos++ + + return p.parseMap() + case '(': + p.pos++ + + x, err := p.parseExpr() + if err != nil { + return nil, err + } + + if err := p.expectCloseParen(); err != nil { + return nil, err + } + + return x, nil + default: + return nil, p.errf("unexpected %q in expression", c) + } +} + +func (p *parser) parseKeyword() (expr, error) { + start := p.pos + + switch word := p.ident(); word { + case "true": + return &literal{value: true}, nil + case "false": + return &literal{value: false}, nil + case "null": + return &literal{value: nil}, nil + default: + p.pos = start + + return nil, p.errf("unexpected %q in expression", word) + } +} + +func (p *parser) parseDoubleQuoted() (expr, error) { + s, err := p.quoted('"') + if err != nil { + return nil, err + } + + if !strings.ContainsAny(s, "$#") { + return &literal{value: s}, nil + } + + sub := &parser{src: s} + + body, term, err := sub.parseBlock() + if err != nil { + return nil, err + } + + if term != "" { + return nil, p.errf("unexpected #%s in string", term) + } + + return &interpolated{body: body}, nil +} + +func (p *parser) parseSingleQuoted() (expr, error) { + s, err := p.quoted('\'') + if err != nil { + return nil, err + } + + return &literal{value: s}, nil +} + +// quoted reads a string delimited by q; a doubled delimiter is an escaped one. +func (p *parser) quoted(q byte) (string, error) { + p.pos++ + + var b strings.Builder + + for p.pos < len(p.src) { + c := p.src[p.pos] + + p.pos++ + + if c != q { + b.WriteByte(c) + + continue + } + + if p.pos < len(p.src) && p.src[p.pos] == q { + b.WriteByte(q) + + p.pos++ + + continue + } + + return b.String(), nil + } + + return "", p.errf("unterminated string") +} + +func (p *parser) parseNumber() (expr, error) { + start := p.pos + + if p.peek() == '-' { + p.pos++ + } + + p.skipDigits() + + isFloat := p.pos+1 < len(p.src) && p.src[p.pos] == '.' && isDigit(p.src[p.pos+1]) + if isFloat { + p.pos++ + p.skipDigits() + } + + text := p.src[start:p.pos] + + if isFloat { + f, err := strconv.ParseFloat(text, 64) + if err != nil { + return nil, p.errf("invalid number %q", text) + } + + return &literal{value: f}, nil + } + + n, err := strconv.ParseInt(text, 10, 64) + if err != nil { + return nil, p.errf("invalid number %q", text) + } + + return &literal{value: n}, nil +} + +func (p *parser) skipDigits() { + for p.pos < len(p.src) && isDigit(p.src[p.pos]) { + p.pos++ + } +} + +func (p *parser) parseListOrRange() (expr, error) { + p.skipSpace() + + if p.peek() == ']' { + p.pos++ + + return &listExpr{}, nil + } + + first, err := p.parseExpr() + if err != nil { + return nil, err + } + + p.skipSpace() + + if p.consume("..") { + to, err := p.parseExpr() + if err != nil { + return nil, err + } + + p.skipSpace() + + if !p.consume("]") { + return nil, p.errf("expected ']' after range") + } + + return &rangeExpr{from: first, to: to}, nil + } + + items := []expr{first} + + if !p.consume("]") { + if !p.consume(",") { + return nil, p.errf("expected ',' or ']'") + } + + rest, err := p.parseArgs(']') + if err != nil { + return nil, err + } + + items = append(items, rest...) + } + + return &listExpr{items: items}, nil +} + +func (p *parser) parseMap() (expr, error) { + m := &mapExpr{} + + p.skipSpace() + + if p.consume("}") { + return m, nil + } + + for { + k, err := p.parseExpr() + if err != nil { + return nil, err + } + + p.skipSpace() + + if !p.consume(":") { + return nil, p.errf("expected ':' in map literal") + } + + v, err := p.parseExpr() + if err != nil { + return nil, err + } + + m.keys = append(m.keys, k) + m.vals = append(m.vals, v) + + p.skipSpace() + + switch { + case p.consume(","): + case p.consume("}"): + return m, nil + default: + return nil, p.errf("expected ',' or '}' in map literal") + } + } +} + +func (p *parser) expectOpenParen() error { + p.skipSpace() + + if !p.consume("(") { + return p.errf("expected '('") + } + + return nil +} + +func (p *parser) expectCloseParen() error { + p.skipSpace() + + if !p.consume(")") { + return p.errf("expected ')'") + } + + return nil +} + +func (p *parser) skipSpace() { + for p.pos < len(p.src) { + switch p.src[p.pos] { + case ' ', '\t', '\n', '\r': + p.pos++ + default: + return + } + } +} + +func (p *parser) peek() byte { + if p.pos >= len(p.src) { + return 0 + } + + return p.src[p.pos] +} + +func (p *parser) consume(tok string) bool { + if strings.HasPrefix(p.src[p.pos:], tok) { + p.pos += len(tok) + + return true + } + + return false +} + +// consumeWord consumes word when it is not followed by an identifier character. +func (p *parser) consumeWord(word string) bool { + end := p.pos + len(word) + if !strings.HasPrefix(p.src[p.pos:], word) || (end < len(p.src) && isIdentChar(p.src[end])) { + return false + } + + p.pos = end + + return true +} + +func (p *parser) ident() string { + start := p.pos + + for p.pos < len(p.src) && isIdentChar(p.src[p.pos]) { + p.pos++ + } + + return p.src[start:p.pos] +} + +func isIdentStart(c byte) bool { return c == '_' || c >= 'a' && c <= 'z' || c >= 'A' && c <= 'Z' } + +// isIdentChar follows Velocity 1.7, which allows '-' inside identifiers. +func isIdentChar(c byte) bool { return isIdentStart(c) || isDigit(c) || c == '-' } + +func isDigit(c byte) bool { return c >= '0' && c <= '9' } diff --git a/internal/vtl/value.go b/internal/vtl/value.go new file mode 100644 index 000000000..6de90ae9b --- /dev/null +++ b/internal/vtl/value.go @@ -0,0 +1,345 @@ +package vtl + +import ( + "bytes" + "encoding/json" + "errors" + "fmt" + "io" + "math" + "sort" + "strconv" + "strings" +) + +// Template values are nil, bool, int64, float64, string, *List, *Map or a host +// Object. Lists and maps are references, so a method such as add or put +// changes the value every variable holding it sees, as in Velocity. + +// Object is a host value a template can read properties from and call methods +// on, such as API Gateway's $input or $util. +type Object interface { + // Get returns the named property. + Get(name string) (any, bool) + // Call invokes the named method. ok is false when the object has no such + // method. + Call(name string, args []any) (result any, ok bool, err error) +} + +// List is an ordered, mutable list value. +type List struct { + Items []any +} + +// NewList returns a list holding items. +func NewList(items ...any) *List { return &List{Items: items} } + +// Index implements jsonpath.Array. +func (l *List) Index(i int) (any, bool) { + if i < 0 || i >= len(l.Items) { + return nil, false + } + + return l.Items[i], true +} + +// Map is an insertion-ordered, mutable map value (Java's LinkedHashMap). +type Map struct { + keys []string + vals map[string]any +} + +// NewMap returns an empty map. +func NewMap() *Map { return &Map{vals: map[string]any{}} } + +// MapOf builds a map from alternating key/value arguments. +func MapOf(kv ...any) *Map { + m := NewMap() + + for i := 0; i+1 < len(kv); i += 2 { + m.Put(fmt.Sprint(kv[i]), kv[i+1]) + } + + return m +} + +// StringMap converts a Go string map to a Map with keys in sorted order. +func StringMap(in map[string]string) *Map { + m := NewMap() + + for _, k := range sortedKeys(in) { + m.Put(k, in[k]) + } + + return m +} + +// Lookup implements jsonpath.Object. +func (m *Map) Lookup(key string) (any, bool) { return m.Get(key) } + +// Get returns the value stored under key. +func (m *Map) Get(key string) (any, bool) { + v, ok := m.vals[key] + + return v, ok +} + +// Put stores v under key and returns the previous value. +func (m *Map) Put(key string, v any) any { + prev, ok := m.vals[key] + if !ok { + m.keys = append(m.keys, key) + } + + m.vals[key] = v + + return prev +} + +// Remove deletes key and returns its previous value. +func (m *Map) Remove(key string) any { + prev, ok := m.vals[key] + if !ok { + return nil + } + + delete(m.vals, key) + + for i, k := range m.keys { + if k == key { + m.keys = append(m.keys[:i], m.keys[i+1:]...) + + break + } + } + + return prev +} + +// Keys returns the keys in insertion order. +func (m *Map) Keys() []string { return append([]string(nil), m.keys...) } + +// Len returns the number of entries. +func (m *Map) Len() int { return len(m.keys) } + +// ParseJSON decodes a JSON document into template values, keeping object key +// order and decoding integral numbers as int64. +func ParseJSON(s string) (any, error) { + dec := json.NewDecoder(strings.NewReader(s)) + dec.UseNumber() + + v, err := decodeValue(dec) + if err != nil { + return nil, err + } + + if _, err := dec.Token(); !errors.Is(err, io.EOF) { + return nil, errorf("unexpected data after JSON value") + } + + return v, nil +} + +func decodeValue(dec *json.Decoder) (any, error) { + tok, err := dec.Token() + if err != nil { + return nil, err + } + + switch t := tok.(type) { + case json.Delim: + if t == '{' { + return decodeObject(dec) + } + + if t == '[' { + return decodeArray(dec) + } + + return nil, errorf("unexpected %q in JSON", t) + case json.Number: + return jsonNumber(t), nil + default: + return t, nil // string, bool or nil + } +} + +func decodeObject(dec *json.Decoder) (any, error) { + m := NewMap() + + for dec.More() { + tok, err := dec.Token() + if err != nil { + return nil, err + } + + key, _ := tok.(string) + + v, err := decodeValue(dec) + if err != nil { + return nil, err + } + + m.Put(key, v) + } + + if _, err := dec.Token(); err != nil { + return nil, err + } + + return m, nil +} + +func decodeArray(dec *json.Decoder) (any, error) { + l := NewList() + + for dec.More() { + v, err := decodeValue(dec) + if err != nil { + return nil, err + } + + l.Items = append(l.Items, v) + } + + if _, err := dec.Token(); err != nil { + return nil, err + } + + return l, nil +} + +func jsonNumber(n json.Number) any { + if i, err := strconv.ParseInt(string(n), 10, 64); err == nil { + return i + } + + f, _ := strconv.ParseFloat(string(n), 64) + + return f +} + +// ToJSON encodes a template value as JSON. Host objects encode as null. +func ToJSON(v any) string { + var b bytes.Buffer + + writeJSON(&b, v) + + return b.String() +} + +func writeJSON(b *bytes.Buffer, v any) { + switch t := v.(type) { + case nil: + b.WriteString("null") + case string: + writeJSONString(b, t) + case bool, int64, float64: + b.WriteString(Stringify(t)) + case *List: + b.WriteByte('[') + + for i, it := range t.Items { + if i > 0 { + b.WriteByte(',') + } + + writeJSON(b, it) + } + + b.WriteByte(']') + case *Map: + b.WriteByte('{') + + for i, k := range t.keys { + if i > 0 { + b.WriteByte(',') + } + + writeJSONString(b, k) + b.WriteByte(':') + writeJSON(b, t.vals[k]) + } + + b.WriteByte('}') + default: + b.WriteString("null") + } +} + +func writeJSONString(b *bytes.Buffer, s string) { + enc := json.NewEncoder(b) + enc.SetEscapeHTML(false) + _ = enc.Encode(s) + + b.Truncate(b.Len() - 1) // drop the encoder's trailing newline +} + +// Stringify renders a value the way Velocity prints it: nil as empty, lists +// as [a, b] and maps as {k=v, k2=v2}, as Java's toString does. +func Stringify(v any) string { + switch t := v.(type) { + case nil: + return "" + case string: + return t + case bool: + return strconv.FormatBool(t) + case int64: + return strconv.FormatInt(t, 10) + case float64: + return formatFloat(t) + default: + return stringifyComposite(v) + } +} + +// stringifyComposite renders lists, maps and Stringers. +func stringifyComposite(v any) string { + switch t := v.(type) { + case *List: + parts := make([]string, len(t.Items)) + for i, it := range t.Items { + parts[i] = Stringify(it) + } + + return "[" + strings.Join(parts, ", ") + "]" + case *Map: + parts := make([]string, len(t.keys)) + for i, k := range t.keys { + parts[i] = k + "=" + Stringify(t.vals[k]) + } + + return "{" + strings.Join(parts, ", ") + "}" + case fmt.Stringer: + return t.String() + default: + return "" + } +} + +// formatFloat prints a double the way Java's Double.toString does for the +// common range: integral values keep a trailing ".0". +func formatFloat(f float64) string { + if math.IsInf(f, 0) || math.IsNaN(f) { + return strconv.FormatFloat(f, 'g', -1, 64) + } + + s := strconv.FormatFloat(f, 'f', -1, 64) + if !strings.Contains(s, ".") { + s += ".0" + } + + return s +} + +func sortedKeys(m map[string]string) []string { + keys := make([]string, 0, len(m)) + for k := range m { + keys = append(keys, k) + } + + sort.Strings(keys) + + return keys +} diff --git a/internal/vtl/vtl_test.go b/internal/vtl/vtl_test.go new file mode 100644 index 000000000..9a303d86a --- /dev/null +++ b/internal/vtl/vtl_test.go @@ -0,0 +1,186 @@ +package vtl + +import ( + "context" + "errors" + "strings" + "testing" + "time" +) + +func render(t *testing.T, src string, vars map[string]any) string { + t.Helper() + + tmpl, err := Parse(src) + if err != nil { + t.Fatalf("Parse(%q): %v", src, err) + } + + res, err := tmpl.Render(context.Background(), vars, RenderOptions{}) + if err != nil { + t.Fatalf("Render(%q): %v", src, err) + } + + return res.Output +} + +func TestRenderDirectivesAndExpressions(t *testing.T) { + body, _ := ParseJSON(`{"name":"pet","tags":["a","b"],"n":3,"nested":{"x":1.5}}`) + + cases := []struct { + name, src, want string + }{ + {"text", "hello", "hello"}, + {"ref", "$name", "pet"}, + {"braced", "${name}s", "pets"}, + {"quiet missing", "[$!missing]", "[]"}, + {"missing renders empty", "[$missing]", "[]"}, + {"dollar not ref", "cost $5", "cost $5"}, + {"escape", `\$name`, "$name"}, + {"property", "$body.nested.x", "1.5"}, + {"index", "$body.tags[1]", "b"}, + {"set and math", "#set($a = 2 + 3 * 4)$a", "14"}, + {"int division", "#set($a = 7 / 2)$a", "3"}, + {"float", "#set($a = 1.5 * 2)$a", "3.0"}, + {"concat", `#set($s = "x" + 1)$s`, "x1"}, + {"interpolated string", `#set($s = "hi $name")$s`, "hi pet"}, + {"single quoted literal", `#set($s = 'hi $name')$s`, "hi $name"}, + {"if else", "#if($body.n > 2)big#else small#end", "big"}, + {"elseif", "#if($body.n == 1)one#elseif($body.n == 3)three#else other#end", "three"}, + {"word ops", "#if($body.n gt 2 and not false)y#end", "y"}, + {"null false", "#if($nope)y#{else}n#end", "n"}, + {"foreach", "#foreach($t in $body.tags)$t$foreach.count#if($foreach.hasNext),#end#end", "a1,b2"}, + {"foreach range", "#foreach($i in [1..3])$i#end", "123"}, + {"foreach map values", `#set($m = {"a": 1, "b": 2})#foreach($v in $m)$v#end`, "12"}, + {"break", "#foreach($i in [1..5])#if($i == 3)#break#end$i#end", "12"}, + {"stop", "a#stop b", "a"}, + {"line comment", "a## comment\nb", "ab"}, + {"block comment", "a#* x *#b", "ab"}, + {"unparsed", "#[[$name]]#", "$name"}, + {"map literal tostring", `#set($m = {"a": 1, "b": "x"})$m`, "{a=1, b=x}"}, + {"list tostring", `#set($l = [1, "two"])$l`, "[1, two]"}, + {"string methods", `$name.toUpperCase() $name.length() $name.substring(1) $name.replace("p", "b")`, "PET 3 et bet"}, + {"string regex", `#set($s = "a1b22")$s.replaceAll("[0-9]+", "-") $s.matches("[a-z0-9]+") $s.split("[0-9]+")`, "a-b- true [a, b]"}, + {"list methods", `#set($l = [])#set($d = $l.add("x"))$l.size() $l.get(0) $l.contains("x")`, "1 x true"}, + {"map methods", `#set($m = {})#set($d = $m.put("k", "v"))$m.get("k") $m.containsKey("k") $m.keySet() $m.size()`, "v true [k] 1"}, + {"set map property", `#set($m = {})#set($m.k = "v")$m`, "{k=v}"}, + {"directive line gobbled", "a\n #set($x = 1)\nb", "a\nb"}, + {"if block gobbled", "#if(true)\n yes\n#end\n", " yes\n"}, + {"hash text", "a # b", "a # b"}, + {"braced directive", "#{if}(true)y#{else}n#{end}", "y"}, + } + + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + got := render(t, c.src, map[string]any{"name": "pet", "body": body}) + if got != c.want { + t.Fatalf("render(%q) = %q, want %q", c.src, got, c.want) + } + }) + } +} + +func TestReturnDirective(t *testing.T) { + tmpl, err := Parse(`before#return({"a": 1})after`) + if err != nil { + t.Fatal(err) + } + + res, err := tmpl.Render(context.Background(), nil, RenderOptions{}) + if err != nil { + t.Fatal(err) + } + + if !res.Returned || res.Output != "before" || ToJSON(res.ReturnValue) != `{"a":1}` { + t.Fatalf("got %+v", res) + } +} + +func TestParseRejectsUnsupported(t *testing.T) { + for _, src := range []string{ + "#macro(x)#end", "#parse('x')", "#include('x')", "#evaluate('x')", "#define($x)#end", + "#if(true)x", "#foreach($i in [1])x", "#end", "#set($x = )", `#set($x = "open)`, + } { + if _, err := Parse(src); err == nil { + t.Errorf("Parse(%q) succeeded, want error", src) + } + } +} + +func TestForeachCapAndStepBudget(t *testing.T) { + src := `#set($l = [])#foreach($i in [1..5000])#set($d = $l.add($i))#end$l.size()` + if got := render(t, src, nil); got != "1000" { + t.Fatalf("foreach cap: got %s", got) + } + + tmpl, _ := Parse(`#foreach($i in [1..1000])#foreach($j in [1..1000])x#end#end`) + + _, err := tmpl.Render(context.Background(), nil, RenderOptions{MaxSteps: 10000}) + if !errors.Is(err, ErrStepBudget) { + t.Fatalf("step budget: err = %v", err) + } + + ctx, cancel := context.WithDeadline(context.Background(), time.Now().Add(-time.Second)) + defer cancel() + + if _, err := tmpl.Render(ctx, nil, RenderOptions{}); !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("deadline: err = %v", err) + } +} + +type hostObj struct{} + +func (hostObj) Get(name string) (any, bool) { + if name == "prop" { + return "P", true + } + + return nil, false +} + +func (hostObj) Call(name string, args []any) (any, bool, error) { + switch name { + case "echo": + return args[0], true, nil + case "fail": + return nil, true, errors.New("boom") + } + + return nil, false, nil +} + +func TestHostObject(t *testing.T) { + if got := render(t, `$h.prop $h.echo("x") [$h.nope()]`, map[string]any{"h": hostObj{}}); got != "P x []" { + t.Fatalf("got %q", got) + } + + tmpl, _ := Parse(`$h.fail()`) + if _, err := tmpl.Render(context.Background(), map[string]any{"h": hostObj{}}, RenderOptions{}); err == nil { + t.Fatal("host error not surfaced") + } +} + +func TestJSONRoundTripKeepsOrder(t *testing.T) { + src := `{"z":1,"a":[true,null,"s",2.5],"m":{"b":2,"a":1}}` + + v, err := ParseJSON(src) + if err != nil { + t.Fatal(err) + } + + if got := ToJSON(v); got != src { + t.Fatalf("ToJSON = %s, want %s", got, src) + } + + if _, err := ParseJSON(`{"a":1} x`); err == nil { + t.Fatal("trailing data accepted") + } + + if got := ToJSON("<&>"); got != `"<&>"` { + t.Fatalf("html escaped: %s", got) + } + + if !strings.Contains(ToJSON(StringMap(map[string]string{"b": "2", "a": "1"})), `{"a":"1","b":"2"}`) { + t.Fatal("StringMap order") + } +} diff --git a/providers/aws/apigateway/dataplane.go b/providers/aws/apigateway/dataplane.go index 09c8f27d8..7bf62177a 100644 --- a/providers/aws/apigateway/dataplane.go +++ b/providers/aws/apigateway/dataplane.go @@ -23,6 +23,7 @@ const ( type resolvedRoute struct { resourceID string resourcePath string + method driver.Method integration driver.Integration pathParameters map[string]string stageVariables map[string]string @@ -57,6 +58,10 @@ func (m *Mock) serveRoute(ctx context.Context, req *driver.ProxyRequest) (*drive return forbiddenMissingToken(), noIntegration } + if route.integration.Type == driver.IntegrationMock { + return m.serveMock(ctx, req, &route), 0 + } + if !isLambdaProxy(route.integration.Type) { return jsonResponse(statusBadGway, `{"message": "Internal server error"}`), noIntegration } @@ -107,10 +112,13 @@ func (m *Mock) resolve(req *driver.ProxyRequest) (resolvedRoute, bool) { return resolvedRoute{}, false } + method := copyMethod(match.method) + return resolvedRoute{ resourceID: match.resource.ID, resourcePath: match.resource.Path, - integration: *match.method.Integration, + method: method, + integration: *method.Integration, pathParameters: match.pathParameters, stageVariables: copyStrMap(st.Variables), apiID: req.RestAPIID, diff --git a/providers/aws/apigateway/mapping.go b/providers/aws/apigateway/mapping.go new file mode 100644 index 000000000..65c255580 --- /dev/null +++ b/providers/aws/apigateway/mapping.go @@ -0,0 +1,356 @@ +package apigateway + +import ( + "context" + "encoding/base64" + "fmt" + "net/url" + "strings" + "time" + "unicode" + "unicode/utf16" + + "github.com/stackshy/cloudemu/v2/internal/jsonpath" + "github.com/stackshy/cloudemu/v2/internal/vtl" + "github.com/stackshy/cloudemu/v2/services/apigateway/driver" +) + +// templateTimeout bounds one mapping-template render. +const templateTimeout = 2 * time.Second + +// $input.params() location keys and the $input.body property. +const ( + locPath = "path" + locQuery = "querystring" + locHeader = "header" + propBody = "body" +) + +// requestTimeLayout is $context.requestTime's CLF format. +const requestTimeLayout = "02/Jan/2006:15:04:05 -0700" + +// mappingContext is the per-request state the $input, $context, +// $stageVariables and $util variables are built from. +type mappingContext struct { + req *driver.ProxyRequest + route *resolvedRoute + account string + reqID string + now time.Time + // context is the $context map. It is shared by the request and response + // templates, so $context.responseOverride set in either survives. + context *vtl.Map +} + +func newMappingContext(req *driver.ProxyRequest, route *resolvedRoute, account, reqID string, now time.Time) *mappingContext { + mc := &mappingContext{req: req, route: route, account: account, reqID: reqID, now: now} + mc.context = mc.buildContext() + + return mc +} + +func (mc *mappingContext) buildContext() *vtl.Map { + identity := vtl.NewMap() + identity.Put("sourceIp", mc.req.SourceIP) + identity.Put("userAgent", headerValue(mc.req.Headers, "User-Agent")) + + override := vtl.NewMap() + override.Put(locHeader, vtl.NewMap()) + + ctx := vtl.NewMap() + for _, kv := range [][2]string{ + {"accountId", mc.account}, + {"apiId", mc.route.apiID}, + {"domainName", mc.req.Host}, + {"domainPrefix", strings.SplitN(mc.req.Host, ".", 2)[0]}, //nolint:mnd // first DNS label + {"extendedRequestId", mc.reqID}, + {"httpMethod", mc.req.HTTPMethod}, + {locPath, "/" + mc.req.StageName + mc.req.Path}, + {"protocol", orDefault(mc.req.Protocol, "HTTP/1.1")}, + {"requestId", mc.reqID}, + {"requestTime", mc.now.UTC().Format(requestTimeLayout)}, + {"resourceId", mc.route.resourceID}, + {"resourcePath", mc.route.resourcePath}, + {"stage", mc.req.StageName}, + } { + ctx.Put(kv[0], kv[1]) + } + + ctx.Put("requestTimeEpoch", mc.now.UnixMilli()) + ctx.Put("identity", identity) + ctx.Put("responseOverride", override) + + return ctx +} + +// render evaluates a mapping template with body as $input's payload. +func (mc *mappingContext) render(ctx context.Context, src, body string) (string, error) { + tmpl, err := vtl.Parse(src) + if err != nil { + return "", err + } + + ctx, cancel := context.WithTimeout(ctx, templateTimeout) + defer cancel() + + vars := map[string]any{ + "input": &inputObject{body: body, params: mc.params()}, + "context": mc.context, + "stageVariables": vtl.StringMap(mc.route.stageVariables), + "util": utilObject{}, + } + + res, err := tmpl.Render(ctx, vars, vtl.RenderOptions{}) + if err != nil { + return "", err + } + + return res.Output, nil +} + +// params is $input.params(): the request's path, querystring and header maps. +func (mc *mappingContext) params() *vtl.Map { + p := vtl.NewMap() + p.Put(locPath, vtl.StringMap(mc.route.pathParameters)) + p.Put(locQuery, vtl.StringMap(mc.req.Query)) + p.Put(locHeader, vtl.StringMap(mc.req.Headers)) + + return p +} + +// responseOverride returns the status and headers a template set through +// $context.responseOverride. +func (mc *mappingContext) responseOverride() (status int, headers map[string]string) { + ov, _ := mc.context.Get("responseOverride") + + om, ok := ov.(*vtl.Map) + if !ok { + return 0, nil + } + + if s, ok := om.Get("status"); ok { + switch v := s.(type) { + case int64: + status = int(v) + case string: + _, _ = fmt.Sscanf(v, "%d", &status) + } + } + + if hv, ok := om.Get(locHeader); ok { + if hm, ok := hv.(*vtl.Map); ok && hm.Len() > 0 { + headers = map[string]string{} + + for _, k := range hm.Keys() { + v, _ := hm.Get(k) + headers[k] = vtl.Stringify(v) + } + } + } + + return status, headers +} + +// inputObject is $input. +type inputObject struct { + body string + params *vtl.Map + parsed any + done bool +} + +func (in *inputObject) Get(name string) (any, bool) { + if name == propBody { + return in.body, true + } + + return nil, false +} + +func (in *inputObject) Call(name string, args []any) (res any, found bool, callErr error) { + switch name { + case locPath: + v, err := in.path(stringArg(args)) + + return v, true, err + case "json": + v, err := in.path(stringArg(args)) + if err != nil { + return nil, true, err + } + + return vtl.ToJSON(v), true, nil + case "params": + if len(args) == 0 { + return in.params, true, nil + } + + return in.param(stringArg(args)), true, nil + case propBody: + return in.body, true, nil + } + + return nil, false, nil +} + +// path evaluates a JSONPath against the JSON body. An empty body is treated as +// an empty object, as API Gateway does. +func (in *inputObject) path(p string) (any, error) { + if !in.done { + in.done = true + + if strings.TrimSpace(in.body) == "" { + in.parsed = vtl.NewMap() + } else if v, err := vtl.ParseJSON(in.body); err == nil { + in.parsed = v + } else { + in.parsed = in.body + } + } + + v, _, err := jsonpath.Eval(p, in.parsed) + + return v, err +} + +// param looks a name up in the path, querystring and header maps, in that +// order, and returns an empty string when it is absent. +func (in *inputObject) param(name string) any { + for _, loc := range []string{locPath, locQuery, locHeader} { + m, _ := in.params.Get(loc) + + lm, ok := m.(*vtl.Map) + if !ok { + continue + } + + if loc == locHeader { + for _, k := range lm.Keys() { + if strings.EqualFold(k, name) { + v, _ := lm.Get(k) + + return v + } + } + + continue + } + + if v, ok := lm.Get(name); ok { + return v + } + } + + return "" +} + +// utilObject is $util. +type utilObject struct{} + +func (utilObject) Get(string) (any, bool) { return nil, false } + +func (utilObject) Call(name string, args []any) (res any, found bool, callErr error) { + s := stringArg(args) + + switch name { + case "escapeJavaScript": + return escapeJavaScript(s), true, nil + case "parseJson": + v, err := vtl.ParseJSON(s) + if err != nil { + return nil, true, fmt.Errorf("$util.parseJson: %w", err) + } + + return v, true, nil + case "urlEncode": + return url.QueryEscape(s), true, nil + case "urlDecode": + v, err := url.QueryUnescape(s) + if err != nil { + return nil, true, fmt.Errorf("$util.urlDecode: %w", err) + } + + return v, true, nil + case "base64Encode": + return base64.StdEncoding.EncodeToString([]byte(s)), true, nil + case "base64Decode": + b, err := base64.StdEncoding.DecodeString(s) + if err != nil { + return nil, true, fmt.Errorf("$util.base64Decode: %w", err) + } + + return string(b), true, nil + } + + return nil, false, nil +} + +func stringArg(args []any) string { + if len(args) == 0 { + return "" + } + + return vtl.Stringify(args[0]) +} + +// escapeJavaScript matches Apache Commons StringEscapeUtils.escapeJavaScript, +// which $util.escapeJavaScript uses: quotes, backslash and '/' are escaped, +// control characters use their short or \uXXXX form, and non-ASCII characters +// become \uXXXX. +func escapeJavaScript(s string) string { + var b strings.Builder + + for _, r := range s { + switch r { + case '"', '\'', '\\', '/': + b.WriteByte('\\') + b.WriteRune(r) + case '\b': + b.WriteString(`\b`) + case '\f': + b.WriteString(`\f`) + case '\n': + b.WriteString(`\n`) + case '\r': + b.WriteString(`\r`) + case '\t': + b.WriteString(`\t`) + default: + writeEscapedRune(&b, r) + } + } + + return b.String() +} + +func writeEscapedRune(b *strings.Builder, r rune) { + const ( + firstPrintable = 0x20 + lastASCII = 0x7f + ) + + if r >= firstPrintable && r < lastASCII { + b.WriteRune(r) + + return + } + + if hi, lo := utf16.EncodeRune(r); hi != unicode.ReplacementChar { + fmt.Fprintf(b, `\u%04X\u%04X`, hi, lo) + + return + } + + fmt.Fprintf(b, `\u%04X`, r) +} + +// headerValue looks a header up case-insensitively. +func headerValue(headers map[string]string, name string) string { + for k, v := range headers { + if strings.EqualFold(k, name) { + return v + } + } + + return "" +} diff --git a/providers/aws/apigateway/mapping_validation.go b/providers/aws/apigateway/mapping_validation.go new file mode 100644 index 000000000..616fbdb98 --- /dev/null +++ b/providers/aws/apigateway/mapping_validation.go @@ -0,0 +1,159 @@ +package apigateway + +import ( + "regexp" + "sort" + "strings" + + cerrors "github.com/stackshy/cloudemu/v2/errors" + "github.com/stackshy/cloudemu/v2/services/apigateway/driver" +) + +// Parameter-mapping expressions API Gateway accepts. +var ( + methodRequestParamKey = regexp.MustCompile(`^method\.request\.` + paramLocations + `\.\S+$`) + methodResponseParamKey = regexp.MustCompile(`^method\.response\.header\.\S+$`) + integrationRequestKey = regexp.MustCompile(`^integration\.request\.` + paramLocations + `\.\S+$`) + integrationRespSource = regexp.MustCompile( + `^integration\.response\.((header|multivalueheader)\.\S+|body(\..+)?)$`) + commonSource = regexp.MustCompile(`^('[^']*'|stageVariables\.\S+|context\.\S+)$`) + methodBody = regexp.MustCompile(`^method\.request\.body(\..+)?$`) +) + +// paramLocations are the request parameter locations a mapping can name. +const paramLocations = `(path|querystring|multivaluequerystring|header|multivalueheader)` + +const ( + mappingErrPrefix = "Invalid mapping expression specified: Validation Result: warnings : [], errors : [" + msgBadThroughBehavior = "Invalid passthrough behavior specified" + msgInvalidSelection = "Invalid selection pattern specified" + msgContentHandlingEnum = "1 validation error detected: Value '%s' at 'contentHandling' failed to satisfy " + + "constraint: Member must satisfy enum value set: [CONVERT_TO_BINARY, CONVERT_TO_TEXT]" +) + +func invalidExpression(expr string) error { + return cerrors.New(cerrors.InvalidArgument, mappingErrPrefix+"Invalid mapping expression specified: "+expr+"]") +} + +func invalidParameter(param string) error { + return cerrors.New(cerrors.InvalidArgument, mappingErrPrefix+"Invalid mapping expression parameter specified: "+param+"]") +} + +// validateMethodRequestParams checks method.request.{location}.{name} keys. +func validateMethodRequestParams(params map[string]bool) error { + for _, k := range sortedBoolKeys(params) { + if !methodRequestParamKey.MatchString(k) { + return invalidExpression(k) + } + } + + return nil +} + +// validateMethodResponseParams checks method.response.header.{name} keys. +func validateMethodResponseParams(params map[string]bool) error { + for _, k := range sortedBoolKeys(params) { + if !methodResponseParamKey.MatchString(k) { + return invalidExpression(k) + } + } + + return nil +} + +func validateContentHandling(v string) error { + switch v { + case "", "CONVERT_TO_BINARY", "CONVERT_TO_TEXT": + return nil + default: + return cerrors.Newf(cerrors.InvalidArgument, msgContentHandlingEnum, v) + } +} + +// validateIntegrationSettings checks an integration's passthrough behavior, +// content handling and request parameter mappings. A method.request source +// must be declared on the method. +func validateIntegrationSettings(ig *driver.Integration, methodParams map[string]bool) error { + switch ig.PassthroughBehavior { + case driver.PassthroughWhenNoMatch, driver.PassthroughWhenNoTemplates, driver.PassthroughNever: + default: + return cerrors.New(cerrors.InvalidArgument, msgBadThroughBehavior) + } + + if err := validateContentHandling(ig.ContentHandling); err != nil { + return err + } + + for _, k := range sortedKeys(ig.RequestParameters) { + if !integrationRequestKey.MatchString(k) { + return invalidParameter(k) + } + + src := ig.RequestParameters[k] + + switch { + case commonSource.MatchString(src), methodBody.MatchString(src): + case methodRequestParamKey.MatchString(src): + if _, ok := methodParams[src]; !ok { + return invalidParameter(src) + } + default: + return invalidExpression(src) + } + } + + return nil +} + +// validateIntegrationResponse checks the selection pattern and that every +// response parameter targets a header the method response declares. +func validateIntegrationResponse(ir *driver.IntegrationResponse, mr *driver.MethodResponse) error { + if _, err := regexp.Compile(ir.SelectionPattern); err != nil { + return cerrors.New(cerrors.InvalidArgument, msgInvalidSelection) + } + + for _, k := range sortedKeys(ir.ResponseParameters) { + declared := mr != nil && mr.ResponseParameters != nil + if declared { + _, declared = mr.ResponseParameters[k] + } + + if !methodResponseParamKey.MatchString(k) || !declared { + return invalidParameter(k) + } + + src := ir.ResponseParameters[k] + if !commonSource.MatchString(src) && !integrationRespSource.MatchString(src) { + return invalidExpression(src) + } + } + + return nil +} + +func sortedBoolKeys(m map[string]bool) []string { + out := make([]string, 0, len(m)) + for k := range m { + out = append(out, k) + } + + sort.Strings(out) + + return out +} + +func sortedKeys(m map[string]string) []string { + out := make([]string, 0, len(m)) + for k := range m { + out = append(out, k) + } + + sort.Strings(out) + + return out +} + +// headerName returns the {name} of a method.response.header.{name} key. +func headerName(key string) string { + return strings.TrimPrefix(key, "method.response.header.") +} diff --git a/providers/aws/apigateway/methods.go b/providers/aws/apigateway/methods.go index 844b14eff..de46c8039 100644 --- a/providers/aws/apigateway/methods.go +++ b/providers/aws/apigateway/methods.go @@ -24,6 +24,10 @@ func (m *Mock) PutMethod( return nil, cerrors.New(cerrors.NotFound, msgResourceNotFound) } + if err := validateMethodRequestParams(in.RequestParameters); err != nil { + return nil, err + } + method := normalizeMethod(httpMethod) if !validHTTPMethod(method) { return nil, cerrors.New(cerrors.InvalidArgument, msgInvalidHTTPMethod) @@ -37,10 +41,13 @@ func (m *Mock) PutMethod( HTTPMethod: method, AuthorizationType: orDefault(in.AuthorizationType, "NONE"), APIKeyRequired: in.APIKeyRequired, + OperationName: in.OperationName, + RequestParameters: copyBoolMap(in.RequestParameters), + RequestModels: copyStrMap(in.RequestModels), } res.Methods[method] = mth - out := *mth + out := copyMethod(mth) return &out, nil } @@ -114,12 +121,23 @@ func (m *Mock) PutIntegration( Type: in.Type, IntegrationHTTPMethod: in.IntegrationHTTPMethod, URI: in.URI, - PassthroughBehavior: orDefault(in.PassthroughBehavior, "WHEN_NO_MATCH"), + PassthroughBehavior: orDefault(in.PassthroughBehavior, driver.PassthroughWhenNoMatch), TimeoutInMillis: orDefaultInt(in.TimeoutInMillis, defaultIntegrationTimeoutMillis), + Credentials: in.Credentials, + RequestParameters: copyStrMap(in.RequestParameters), + RequestTemplates: copyStrMap(in.RequestTemplates), + ContentHandling: in.ContentHandling, + CacheNamespace: orDefault(in.CacheNamespace, resourceID), + CacheKeyParameters: append([]string(nil), in.CacheKeyParameters...), } + + if err := validateIntegrationSettings(ig, mth.RequestParameters); err != nil { + return nil, err + } + mth.Integration = ig - out := *ig + out := copyIntegration(ig) return &out, nil } @@ -190,12 +208,7 @@ func (m *Mock) lookupMethod(restAPIID, resourceID, httpMethod string) (*driver.M return nil, cerrors.New(cerrors.NotFound, msgMethodNotFound) } - out := *mth - - if mth.Integration != nil { - ig := *mth.Integration - out.Integration = &ig - } + out := copyMethod(mth) return &out, nil } diff --git a/providers/aws/apigateway/mock_integration.go b/providers/aws/apigateway/mock_integration.go new file mode 100644 index 000000000..0203298d3 --- /dev/null +++ b/providers/aws/apigateway/mock_integration.go @@ -0,0 +1,347 @@ +package apigateway + +import ( + "context" + "regexp" + "strconv" + "strings" + + "github.com/stackshy/cloudemu/v2/internal/idgen" + "github.com/stackshy/cloudemu/v2/internal/jsonpath" + "github.com/stackshy/cloudemu/v2/internal/vtl" + "github.com/stackshy/cloudemu/v2/services/apigateway/driver" +) + +// Gateway error statuses and bodies for mapping failures. +const ( + statusOK = 200 + statusInternalServerError = 500 + statusUnsupportedMediaType = 415 + bodyInternalServerError = `{"message": "Internal server error"}` + bodyUnsupportedMediaType = `{"message": "Unsupported Media Type"}` + contentTypeJSON = "application/json" + headerContentType = "Content-Type" + headerRequestID = "x-amzn-RequestId" + headerErrorType = "x-amzn-ErrorType" +) + +// serveMock answers a MOCK integration: the request template (picked by +// Content-Type, subject to passthroughBehavior) yields a JSON document whose +// statusCode selects the integration response, whose template renders the +// body. +func (m *Mock) serveMock(ctx context.Context, req *driver.ProxyRequest, route *resolvedRoute) *driver.ProxyResponse { + mc := newMappingContext(req, route, m.opts.AccountID, idgen.UUID(), m.opts.Clock.Now()) + + payload, templated, rejected := mapRequest(ctx, mc, &route.integration, req) + if rejected != nil { + return rejected + } + + status, ok := mockStatusCode(payload, templated) + if !ok { + return gatewayError(statusInternalServerError, bodyInternalServerError, "InternalServerErrorException", mc.reqID) + } + + return mapResponse(ctx, mc, route, strconv.Itoa(status), "", nil) +} + +// mapRequest applies the integration request template. It returns the +// integration payload and whether a template produced it, or a 415/500 +// response when the request is rejected. +func mapRequest( + ctx context.Context, mc *mappingContext, ig *driver.Integration, req *driver.ProxyRequest, +) (payload string, templated bool, rejected *driver.ProxyResponse) { + ct := mediaType(headerValue(req.Headers, headerContentType)) + if ct == "" { + ct = contentTypeJSON + } + + tmpl, found := lookupTemplate(ig.RequestTemplates, ct) + if !found { + if passthroughAllowed(ig) { + return req.Body, false, nil + } + + return "", false, gatewayError(statusUnsupportedMediaType, bodyUnsupportedMediaType, + "UnsupportedMediaTypeException", mc.reqID) + } + + out, err := mc.render(ctx, tmpl, req.Body) + if err != nil { + return "", false, gatewayError(statusInternalServerError, bodyInternalServerError, + "InternalServerErrorException", mc.reqID) + } + + return out, true, nil +} + +// passthroughAllowed applies passthroughBehavior to a request whose content +// type has no template. +func passthroughAllowed(ig *driver.Integration) bool { + switch ig.PassthroughBehavior { + case driver.PassthroughNever: + return false + case driver.PassthroughWhenNoTemplates: + return len(ig.RequestTemplates) == 0 + default: + return true + } +} + +// mockStatusCode reads statusCode from a MOCK request payload. A payload that +// is not JSON is a configuration error only when a template produced it; a +// missing statusCode means 200. +func mockStatusCode(payload string, templated bool) (int, bool) { + if strings.TrimSpace(payload) == "" { + return statusOK, true + } + + v, err := vtl.ParseJSON(payload) + if err != nil { + return statusOK, !templated + } + + doc, ok := v.(*vtl.Map) + if !ok { + return statusOK, true + } + + raw, ok := doc.Get("statusCode") + if !ok { + return statusOK, true + } + + switch n := raw.(type) { + case int64: + return int(n), true + case float64: + if n == float64(int(n)) { + return int(n), true + } + } + + return 0, false +} + +// mapResponse selects the integration response for backendStatus, renders its +// template against the backend body and applies its header mappings and any +// $context.responseOverride. With no matching and no default integration +// response the request fails with a 500, as API Gateway does. +func mapResponse( + ctx context.Context, mc *mappingContext, route *resolvedRoute, + backendStatus, backendBody string, backendHeaders map[string]string, +) *driver.ProxyResponse { + ir := selectIntegrationResponse(route.integration.IntegrationResponses, backendStatus) + if ir == nil { + return gatewayError(statusInternalServerError, bodyInternalServerError, "InternalServerErrorException", mc.reqID) + } + + status, err := strconv.Atoi(ir.StatusCode) + if err != nil { + return gatewayError(statusInternalServerError, bodyInternalServerError, "InternalServerErrorException", mc.reqID) + } + + resp := &driver.ProxyResponse{ + StatusCode: status, + Headers: map[string]string{headerContentType: contentTypeJSON, headerRequestID: mc.reqID}, + Body: backendBody, + } + + if ct, tmpl, ok := selectResponseTemplate(ir.ResponseTemplates, headerValue(mc.req.Headers, "Accept")); ok { + resp.Headers[headerContentType] = ct + + if strings.TrimSpace(tmpl) != "" { + out, err := mc.render(ctx, tmpl, backendBody) + if err != nil { + return gatewayError(statusInternalServerError, bodyInternalServerError, "InternalServerErrorException", mc.reqID) + } + + resp.Body = out + } + } + + for _, key := range sortedKeys(ir.ResponseParameters) { + if v, ok := mc.resolveResponseSource(ir.ResponseParameters[key], backendBody, backendHeaders); ok { + resp.Headers[headerName(key)] = v + } + } + + overrideStatus, overrideHeaders := mc.responseOverride() + if overrideStatus != 0 { + resp.StatusCode = overrideStatus + } + + for k, v := range overrideHeaders { + resp.Headers[k] = v + } + + return resp +} + +// selectIntegrationResponse returns the first integration response (by status +// code) whose selection pattern fully matches match, else the default one +// (empty pattern), else nil. +func selectIntegrationResponse(irs map[string]*driver.IntegrationResponse, match string) *driver.IntegrationResponse { + var def *driver.IntegrationResponse + + for _, code := range sortedResponseCodes(irs) { + ir := irs[code] + + if ir.SelectionPattern == "" { + if def == nil { + def = ir + } + + continue + } + + re, err := regexp.Compile("^(?:" + ir.SelectionPattern + ")$") + if err == nil && re.MatchString(match) { + return ir + } + } + + return def +} + +// selectResponseTemplate picks a response template by the request's Accept +// header, falling back to application/json and then to the first content type. +func selectResponseTemplate(templates map[string]string, accept string) (contentType, tmpl string, ok bool) { + if len(templates) == 0 { + return "", "", false + } + + for _, part := range strings.Split(accept, ",") { + if ct := mediaType(part); ct != "" { + if key, found := templateKey(templates, ct); found { + return key, templates[key], true + } + } + } + + if key, found := templateKey(templates, contentTypeJSON); found { + return key, templates[key], true + } + + key := sortedKeys(templates)[0] + + return key, templates[key], true +} + +// lookupTemplate finds the template for a content type, case-insensitively. +func lookupTemplate(templates map[string]string, ct string) (string, bool) { + key, ok := templateKey(templates, ct) + if !ok { + return "", false + } + + return templates[key], true +} + +func templateKey(templates map[string]string, ct string) (string, bool) { + for _, k := range sortedKeys(templates) { + if strings.EqualFold(k, ct) { + return k, true + } + } + + return "", false +} + +// mediaType strips parameters and whitespace from a Content-Type or Accept +// entry and lower-cases it. +func mediaType(v string) string { + if i := strings.IndexByte(v, ';'); i >= 0 { + v = v[:i] + } + + return strings.ToLower(strings.TrimSpace(v)) +} + +// resolveResponseSource evaluates an integration response parameter source: +// a 'static' value, an integration.response header or body, a stage variable +// or a $context value. +func (mc *mappingContext) resolveResponseSource(src, body string, headers map[string]string) (string, bool) { + const ( + headerPrefix = "integration.response.header." + multiHeaderPref = "integration.response.multivalueheader." + bodyRef = "integration.response.body" + stagePrefix = "stageVariables." + contextPrefix = "context." + ) + + switch { + case len(src) >= 2 && src[0] == '\'' && src[len(src)-1] == '\'': + return src[1 : len(src)-1], true + case strings.HasPrefix(src, headerPrefix): + v := headerValue(headers, strings.TrimPrefix(src, headerPrefix)) + + return v, v != "" + case strings.HasPrefix(src, multiHeaderPref): + v := headerValue(headers, strings.TrimPrefix(src, multiHeaderPref)) + + return v, v != "" + case src == bodyRef: + return body, true + case strings.HasPrefix(src, bodyRef+"."): + return bodyPath(body, "$"+strings.TrimPrefix(src, bodyRef)) + case strings.HasPrefix(src, stagePrefix): + v, ok := mc.route.stageVariables[strings.TrimPrefix(src, stagePrefix)] + + return v, ok + case strings.HasPrefix(src, contextPrefix): + return contextValue(mc.context, strings.TrimPrefix(src, contextPrefix)) + } + + return "", false +} + +func bodyPath(body, path string) (string, bool) { + doc, err := vtl.ParseJSON(body) + if err != nil { + return "", false + } + + v, ok, err := jsonpath.Eval(path, doc) + if err != nil || !ok || v == nil { + return "", false + } + + switch v.(type) { + case *vtl.Map, *vtl.List: + return vtl.ToJSON(v), true + default: + return vtl.Stringify(v), true + } +} + +// contextValue walks a dotted path (e.g. identity.sourceIp) into $context. +func contextValue(ctx *vtl.Map, path string) (string, bool) { + var cur any = ctx + + for _, part := range strings.Split(path, ".") { + m, ok := cur.(*vtl.Map) + if !ok { + return "", false + } + + if cur, ok = m.Get(part); !ok { + return "", false + } + } + + s := vtl.Stringify(cur) + + return s, s != "" +} + +// gatewayError is an error API Gateway itself produces. +func gatewayError(status int, body, errType, reqID string) *driver.ProxyResponse { + return &driver.ProxyResponse{ + StatusCode: status, + Headers: map[string]string{ + headerContentType: contentTypeJSON, headerErrorType: errType, headerRequestID: reqID, + }, + Body: body, + } +} diff --git a/providers/aws/apigateway/mock_integration_test.go b/providers/aws/apigateway/mock_integration_test.go new file mode 100644 index 000000000..7185ef907 --- /dev/null +++ b/providers/aws/apigateway/mock_integration_test.go @@ -0,0 +1,267 @@ +package apigateway_test + +import ( + "testing" + + "github.com/stackshy/cloudemu/v2/errors" + "github.com/stackshy/cloudemu/v2/providers/aws/apigateway" + "github.com/stackshy/cloudemu/v2/services/apigateway/driver" +) + +// mockMethod builds /m with a MOCK GET using the given request templates and +// passthrough behaviour, a 200 method response declaring X-A, and a default +// integration response with responseTemplate. It deploys stage "s" with +// stage variable v=sv and returns the API and resource ids. +func mockMethod( + t *testing.T, m *apigateway.Mock, reqTemplates map[string]string, passthrough, responseTemplate string, +) (apiID, resID string) { + t.Helper() + + api, err := m.CreateRestAPI(ctx(), &driver.CreateRestAPIInput{Name: "mock"}) + if err != nil { + t.Fatalf("CreateRestAPI: %v", err) + } + + res, err := m.CreateResource(ctx(), api.ID, api.RootResourceID, "m") + if err != nil { + t.Fatalf("CreateResource: %v", err) + } + + if _, err := m.PutMethod(ctx(), api.ID, res.ID, "GET", driver.PutMethodInput{}); err != nil { + t.Fatalf("PutMethod: %v", err) + } + + if _, err := m.PutIntegration(ctx(), api.ID, res.ID, "GET", driver.PutIntegrationInput{ + Type: driver.IntegrationMock, RequestTemplates: reqTemplates, PassthroughBehavior: passthrough, + }); err != nil { + t.Fatalf("PutIntegration: %v", err) + } + + if _, err := m.PutMethodResponse(ctx(), api.ID, res.ID, "GET", "200", driver.PutMethodResponseInput{ + ResponseParameters: map[string]bool{"method.response.header.X-A": false}, + }); err != nil { + t.Fatalf("PutMethodResponse: %v", err) + } + + in := driver.PutIntegrationResponseInput{ + ResponseParameters: map[string]string{"method.response.header.X-A": "stageVariables.v"}, + } + if responseTemplate != "" { + in.ResponseTemplates = map[string]string{"application/json": responseTemplate} + } + + if _, err := m.PutIntegrationResponse(ctx(), api.ID, res.ID, "GET", "200", in); err != nil { + t.Fatalf("PutIntegrationResponse: %v", err) + } + + if _, err := m.CreateDeployment(ctx(), api.ID, driver.CreateDeploymentInput{ + StageName: "s", Variables: map[string]string{"v": "sv"}, + }); err != nil { + t.Fatalf("CreateDeployment: %v", err) + } + + return api.ID, res.ID +} + +func invokeMock(t *testing.T, m *apigateway.Mock, apiID string, req driver.ProxyRequest) *driver.ProxyResponse { + t.Helper() + + req.RestAPIID, req.StageName, req.HTTPMethod, req.Path = apiID, "s", "GET", "/m" + + resp, err := m.InvokeRoute(ctx(), &req) + if err != nil { + t.Fatalf("InvokeRoute: %v", err) + } + + return resp +} + +func TestMockResponseTemplateContext(t *testing.T) { + m := newMock(t) + tmpl := `{"q":"$input.params('q')","h":"$input.params('x-h')","path":"$context.resourcePath",` + + `"esc":"$util.escapeJavaScript('a"b/é')","b64":"$util.base64Encode('hi')",` + + `"url":"$util.urlEncode('a b')","json":$input.json('$')}` + apiID, _ := mockMethod(t, m, map[string]string{"application/json": `{"statusCode": 200}`}, "", tmpl) + + resp := invokeMock(t, m, apiID, driver.ProxyRequest{ + Query: map[string]string{"q": "1"}, Headers: map[string]string{"X-H": "hv"}, + }) + + want := `{"q":"1","h":"hv","path":"/m","esc":"a\"b\/\u00E9","b64":"aGk=","url":"a+b","json":{}}` + if resp.StatusCode != 200 || resp.Body != want || resp.Headers["X-A"] != "sv" { + t.Fatalf("got %d %s %v\nwant %s", resp.StatusCode, resp.Body, resp.Headers, want) + } +} + +func TestMockResponseOverride(t *testing.T) { + m := newMock(t) + tmpl := `#set($context.responseOverride.status = 201)#set($context.responseOverride.header.X-O = "o")created` + apiID, _ := mockMethod(t, m, map[string]string{"application/json": `{"statusCode": 200}`}, "", tmpl) + + resp := invokeMock(t, m, apiID, driver.ProxyRequest{}) + if resp.StatusCode != 201 || resp.Body != "created" || resp.Headers["X-O"] != "o" { + t.Fatalf("got %d %s %v", resp.StatusCode, resp.Body, resp.Headers) + } +} + +func TestMockPassthroughBehavior(t *testing.T) { + cases := []struct { + name string + templates map[string]string + passthrough string + contentType string + wantStatus int + }{ + {"no match passes through", map[string]string{"application/json": `{"statusCode":200}`}, "", "text/plain", 200}, + {"never rejects", map[string]string{"application/json": `{"statusCode":200}`}, driver.PassthroughNever, "text/plain", 415}, + {"no templates passes", nil, driver.PassthroughWhenNoTemplates, "text/plain", 200}, + {"templates defined rejects", map[string]string{"application/json": "{}"}, driver.PassthroughWhenNoTemplates, "text/xml", 415}, + {"missing content type is json", map[string]string{"application/json": `{"statusCode":200}`}, driver.PassthroughNever, "", 200}, + {"bad statusCode", map[string]string{"application/json": `{"statusCode":"x"}`}, "", "", 500}, + {"template not json", map[string]string{"application/json": `not json`}, "", "", 500}, + {"template error", map[string]string{"application/json": `$util.parseJson("{")`}, "", "", 500}, + {"unmatched status", map[string]string{"application/json": `{"statusCode":404}`}, "", "", 200}, + } + + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + m := newMock(t) + apiID, _ := mockMethod(t, m, c.templates, c.passthrough, "") + + resp := invokeMock(t, m, apiID, driver.ProxyRequest{ + Headers: map[string]string{"Content-Type": c.contentType}, Body: "plain", + }) + if resp.StatusCode != c.wantStatus { + t.Fatalf("status = %d, want %d (%s)", resp.StatusCode, c.wantStatus, resp.Body) + } + }) + } +} + +func TestMethodResponseLifecycleAndSnapshot(t *testing.T) { + m := newMock(t) + apiID, resID := mockMethod(t, m, map[string]string{"application/json": `{"statusCode":200}`}, "", "{}") + + _, err := m.PutMethodResponse(ctx(), apiID, resID, "GET", "200", driver.PutMethodResponseInput{}) + assertMessage(t, err, errors.IsAlreadyExists, "Response already exists for this resource") + + _, err = m.PutMethodResponse(ctx(), apiID, resID, "GET", "200", driver.PutMethodResponseInput{ + ResponseParameters: map[string]bool{"bad.key": true}, + }) + if !errors.IsAlreadyExists(err) && !errors.IsInvalidArgument(err) { + t.Fatalf("bad key: %v", err) + } + + mr, err := m.UpdateMethodResponse(ctx(), apiID, resID, "GET", "200", []driver.PatchOperation{ + {Op: "remove", Path: "/responseParameters/method.response.header.X-A"}, + {Op: "add", Path: "/responseModels/application~1json", Value: "Empty"}, + }) + if err != nil || len(mr.ResponseParameters) != 0 || mr.ResponseModels["application/json"] != "Empty" { + t.Fatalf("UpdateMethodResponse = %+v %v", mr, err) + } + + ir, err := m.UpdateIntegrationResponse(ctx(), apiID, resID, "GET", "200", []driver.PatchOperation{ + {Op: "replace", Path: "/selectionPattern", Value: "2\\d\\d"}, + {Op: "remove", Path: "/responseParameters/method.response.header.X-A"}, + {Op: "replace", Path: "/contentHandling", Value: "CONVERT_TO_TEXT"}, + }) + if err != nil || ir.SelectionPattern != `2\d\d` || ir.ContentHandling != "CONVERT_TO_TEXT" { + t.Fatalf("UpdateIntegrationResponse = %+v %v", ir, err) + } + + if _, err := m.UpdateIntegrationResponse(ctx(), apiID, resID, "GET", "200", []driver.PatchOperation{ + {Op: "replace", Path: "/contentHandling", Value: "BOGUS"}, + }); !errors.IsInvalidArgument(err) { + t.Fatalf("bad contentHandling: %v", err) + } + + mth, err := m.UpdateMethod(ctx(), apiID, resID, "GET", []driver.PatchOperation{ + {Op: "add", Path: "/requestParameters/method.request.querystring.q", Value: "true"}, + {Op: "add", Path: "/requestModels/application~1json", Value: "Empty"}, + {Op: "replace", Path: "/operationName", Value: "GetM"}, + }) + if err != nil || !mth.RequestParameters["method.request.querystring.q"] || mth.OperationName != "GetM" { + t.Fatalf("UpdateMethod = %+v %v", mth, err) + } + + ig, err := m.UpdateIntegration(ctx(), apiID, resID, "GET", []driver.PatchOperation{ + {Op: "add", Path: "/requestParameters/integration.request.header.X", Value: "method.request.querystring.q"}, + {Op: "add", Path: "/cacheKeyParameters/method.request.querystring.q"}, + {Op: "replace", Path: "/credentials", Value: "arn:aws:iam::000000000000:role/r"}, + }) + if err != nil || ig.RequestParameters["integration.request.header.X"] == "" || len(ig.CacheKeyParameters) != 1 || + len(ig.IntegrationResponses) != 1 { + t.Fatalf("UpdateIntegration = %+v %v", ig, err) + } + + data, err := m.Snapshot(ctx(), false) + if err != nil { + t.Fatal(err) + } + + dst := newMock(t) + if err := dst.Restore(ctx(), data); err != nil { + t.Fatal(err) + } + + got, err := dst.GetIntegrationResponse(ctx(), apiID, resID, "GET", "200") + if err != nil || got.SelectionPattern != `2\d\d` { + t.Fatalf("restored integration response = %+v %v", got, err) + } + + if mr, err := dst.GetMethodResponse(ctx(), apiID, resID, "GET", "200"); err != nil || mr.ResponseModels["application/json"] != "Empty" { + t.Fatalf("restored method response = %+v %v", mr, err) + } + + // The deployed tree survives too: the restored stage still answers. + if resp := invokeMock(t, dst, apiID, driver.ProxyRequest{}); resp.StatusCode != 200 { + t.Fatalf("restored invoke = %d", resp.StatusCode) + } + + if err := m.DeleteIntegrationResponse(ctx(), apiID, resID, "GET", "200"); err != nil { + t.Fatal(err) + } + + if _, err := m.GetIntegrationResponse(ctx(), apiID, resID, "GET", "200"); !errors.IsNotFound(err) { + t.Fatalf("after delete: %v", err) + } + + if err := m.DeleteMethodResponse(ctx(), apiID, resID, "GET", "200"); err != nil { + t.Fatal(err) + } + + if err := m.DeleteMethodResponse(ctx(), apiID, resID, "GET", "200"); !errors.IsNotFound(err) { + t.Fatalf("double delete: %v", err) + } +} + +func TestMockSelectionPatternPicksResponse(t *testing.T) { + m := newMock(t) + apiID, resID := mockMethod(t, m, map[string]string{"application/json": `{"statusCode": $input.params('c')}`}, "", "ok") + + if _, err := m.PutMethodResponse(ctx(), apiID, resID, "GET", "400", driver.PutMethodResponseInput{}); err != nil { + t.Fatal(err) + } + + if _, err := m.PutIntegrationResponse(ctx(), apiID, resID, "GET", "400", driver.PutIntegrationResponseInput{ + SelectionPattern: `4\d\d`, + ResponseTemplates: map[string]string{"application/json": "bad", "text/plain": "bad-text"}, + }); err != nil { + t.Fatal(err) + } + + if _, err := m.CreateDeployment(ctx(), apiID, driver.CreateDeploymentInput{StageName: "s"}); err != nil { + t.Fatal(err) + } + + resp := invokeMock(t, m, apiID, driver.ProxyRequest{ + Query: map[string]string{"c": "404"}, Headers: map[string]string{"Accept": "text/plain"}, + }) + if resp.StatusCode != 400 || resp.Body != "bad-text" || resp.Headers["Content-Type"] != "text/plain" { + t.Fatalf("4xx = %d %s %v", resp.StatusCode, resp.Body, resp.Headers) + } + + if resp := invokeMock(t, m, apiID, driver.ProxyRequest{Query: map[string]string{"c": "200"}}); resp.Body != "ok" { + t.Fatalf("default = %d %s", resp.StatusCode, resp.Body) + } +} diff --git a/providers/aws/apigateway/patch.go b/providers/aws/apigateway/patch.go index f8b81a3ed..d9f727daf 100644 --- a/providers/aws/apigateway/patch.go +++ b/providers/aws/apigateway/patch.go @@ -17,6 +17,10 @@ const ( opRemove = "remove" ) +// pathContentHandling is the contentHandling patch path shared by +// integrations and integration responses. +const pathContentHandling = "/contentHandling" + // pathDescription is the JSON Pointer for the /description field, shared by the // RestApi, Stage and Deployment patch appliers. const pathDescription = "/description" @@ -182,15 +186,16 @@ func (m *Mock) UpdateMethod( return nil, cerrors.New(cerrors.NotFound, msgMethodNotFound) } + next := copyMethod(mth) for _, op := range ops { - switch op.Path { - case "/authorizationType": - mth.AuthorizationType = op.Value - case "/apiKeyRequired": - mth.APIKeyRequired = parseBool(op.Value) - } + applyMethodPatch(&next, op) + } + + if err := validateMethodRequestParams(next.RequestParameters); err != nil { + return nil, err } + *mth = next out := copyMethod(mth) return &out, nil @@ -218,11 +223,17 @@ func (m *Mock) UpdateIntegration( return nil, cerrors.New(cerrors.NotFound, msgIntegrationNotFound) } + next := copyIntegration(mth.Integration) for _, op := range ops { - applyIntegrationPatch(mth.Integration, op) + applyIntegrationPatch(&next, op) } - out := *mth.Integration + if err := validateIntegrationSettings(&next, mth.RequestParameters); err != nil { + return nil, err + } + + *mth.Integration = next + out := copyIntegration(mth.Integration) return &out, nil } @@ -242,6 +253,51 @@ func applyIntegrationPatch(ig *driver.Integration, op driver.PatchOperation) { if n, err := strconv.Atoi(op.Value); err == nil { ig.TimeoutInMillis = n } + case "/credentials": + ig.Credentials = patchRef(op) + case pathContentHandling: + ig.ContentHandling = patchRef(op) + case "/cacheNamespace": + ig.CacheNamespace = patchRef(op) + default: + applyIntegrationMapPatch(ig, op) + } +} + +// applyIntegrationMapPatch handles the map- and list-valued integration paths. +func applyIntegrationMapPatch(ig *driver.Integration, op driver.PatchOperation) { + applyMapPatch(op, "/requestTemplates/", func(k, v string, remove bool) { + ig.RequestTemplates = patchStrMap(ig.RequestTemplates, k, v, remove) + }) + applyMapPatch(op, "/requestParameters/", func(k, v string, remove bool) { + ig.RequestParameters = patchStrMap(ig.RequestParameters, k, v, remove) + }) + applyMapPatch(op, "/cacheKeyParameters/", func(k, _ string, remove bool) { + action := opAdd + if remove { + action = opRemove + } + + ig.CacheKeyParameters = patchStringSlice(ig.CacheKeyParameters, action, k) + }) +} + +// applyMethodPatch applies one patch op to a Method. +func applyMethodPatch(mth *driver.Method, op driver.PatchOperation) { + switch op.Path { + case "/authorizationType": + mth.AuthorizationType = op.Value + case "/apiKeyRequired": + mth.APIKeyRequired = parseBool(op.Value) + case "/operationName": + mth.OperationName = patchRef(op) + default: + applyMapPatch(op, "/requestParameters/", func(k, v string, remove bool) { + mth.RequestParameters = patchBoolMap(mth.RequestParameters, k, v, remove) + }) + applyMapPatch(op, "/requestModels/", func(k, v string, remove bool) { + mth.RequestModels = patchStrMap(mth.RequestModels, k, v, remove) + }) } } @@ -409,18 +465,6 @@ func patchStringSlice(s []string, op, v string) []string { } } -// copyMethod returns a deep copy of a method and its integration. -func copyMethod(mth *driver.Method) driver.Method { - out := *mth - - if mth.Integration != nil { - ig := *mth.Integration - out.Integration = &ig - } - - return out -} - // parseBool reports whether an on-the-wire patch value (always a string) is // true. func parseBool(v string) bool { diff --git a/providers/aws/apigateway/resources.go b/providers/aws/apigateway/resources.go index ef4a888c2..a6b2f920b 100644 --- a/providers/aws/apigateway/resources.go +++ b/providers/aws/apigateway/resources.go @@ -197,13 +197,7 @@ func copyResource(r *driver.Resource) driver.Resource { out.Methods = make(map[string]*driver.Method, len(r.Methods)) for k, mth := range r.Methods { - cp := *mth - - if mth.Integration != nil { - ig := *mth.Integration - cp.Integration = &ig - } - + cp := copyMethod(mth) out.Methods[k] = &cp } diff --git a/providers/aws/apigateway/responses.go b/providers/aws/apigateway/responses.go new file mode 100644 index 000000000..136a9e03b --- /dev/null +++ b/providers/aws/apigateway/responses.go @@ -0,0 +1,427 @@ +package apigateway + +import ( + "context" + "regexp" + "sort" + + cerrors "github.com/stackshy/cloudemu/v2/errors" + "github.com/stackshy/cloudemu/v2/services/apigateway/driver" +) + +// Error messages for method and integration responses. +const ( + msgResponseNotFound = "Invalid Response status code specified" + msgResponseExists = "Response already exists for this resource" + msgInvalidStatus = "Invalid status code specified" +) + +// responseStatusPattern is the StatusCode shape from the API model. +var responseStatusPattern = regexp.MustCompile(`^[1-5]\d\d$`) + +// PutMethodResponse declares statusCode on a method. +func (m *Mock) PutMethodResponse( + _ context.Context, restAPIID, resourceID, httpMethod, statusCode string, in driver.PutMethodResponseInput, +) (*driver.MethodResponse, error) { + if !responseStatusPattern.MatchString(statusCode) { + return nil, cerrors.New(cerrors.InvalidArgument, msgInvalidStatus) + } + + if err := validateMethodResponseParams(in.ResponseParameters); err != nil { + return nil, err + } + + var out driver.MethodResponse + + err := m.withMethod(restAPIID, resourceID, httpMethod, func(mth *driver.Method) error { + if _, exists := mth.MethodResponses[statusCode]; exists { + return cerrors.New(cerrors.AlreadyExists, msgResponseExists) + } + + mr := &driver.MethodResponse{ + StatusCode: statusCode, + ResponseParameters: copyBoolMap(in.ResponseParameters), + ResponseModels: copyStrMap(in.ResponseModels), + } + + if mth.MethodResponses == nil { + mth.MethodResponses = map[string]*driver.MethodResponse{} + } + + mth.MethodResponses[statusCode] = mr + out = copyMethodResponse(mr) + + return nil + }) + if err != nil { + return nil, err + } + + return &out, nil +} + +// GetMethodResponse returns a declared method response. +func (m *Mock) GetMethodResponse( + _ context.Context, restAPIID, resourceID, httpMethod, statusCode string, +) (*driver.MethodResponse, error) { + mth, err := m.lookupMethod(restAPIID, resourceID, httpMethod) + if err != nil { + return nil, err + } + + mr, ok := mth.MethodResponses[statusCode] + if !ok { + return nil, cerrors.New(cerrors.NotFound, msgResponseNotFound) + } + + return mr, nil +} + +// UpdateMethodResponse patches a method response's parameters and models. +func (m *Mock) UpdateMethodResponse( + _ context.Context, restAPIID, resourceID, httpMethod, statusCode string, ops []driver.PatchOperation, +) (*driver.MethodResponse, error) { + var out driver.MethodResponse + + err := m.withMethod(restAPIID, resourceID, httpMethod, func(mth *driver.Method) error { + mr, ok := mth.MethodResponses[statusCode] + if !ok { + return cerrors.New(cerrors.NotFound, msgResponseNotFound) + } + + next := copyMethodResponse(mr) + + for _, op := range ops { + applyMapPatch(op, "/responseParameters/", func(k, v string, remove bool) { + next.ResponseParameters = patchBoolMap(next.ResponseParameters, k, v, remove) + }) + applyMapPatch(op, "/responseModels/", func(k, v string, remove bool) { + next.ResponseModels = patchStrMap(next.ResponseModels, k, v, remove) + }) + } + + if err := validateMethodResponseParams(next.ResponseParameters); err != nil { + return err + } + + *mr = next + out = copyMethodResponse(mr) + + return nil + }) + if err != nil { + return nil, err + } + + return &out, nil +} + +// DeleteMethodResponse removes a method response. +func (m *Mock) DeleteMethodResponse(_ context.Context, restAPIID, resourceID, httpMethod, statusCode string) error { + return m.withMethod(restAPIID, resourceID, httpMethod, func(mth *driver.Method) error { + if _, ok := mth.MethodResponses[statusCode]; !ok { + return cerrors.New(cerrors.NotFound, msgResponseNotFound) + } + + delete(mth.MethodResponses, statusCode) + + return nil + }) +} + +// PutIntegrationResponse creates or replaces the integration response for +// statusCode. The method must already declare that status code. +func (m *Mock) PutIntegrationResponse( + _ context.Context, restAPIID, resourceID, httpMethod, statusCode string, in driver.PutIntegrationResponseInput, +) (*driver.IntegrationResponse, error) { + if !responseStatusPattern.MatchString(statusCode) { + return nil, cerrors.New(cerrors.InvalidArgument, msgInvalidStatus) + } + + if err := validateContentHandling(in.ContentHandling); err != nil { + return nil, err + } + + var out driver.IntegrationResponse + + err := m.withMethod(restAPIID, resourceID, httpMethod, func(mth *driver.Method) error { + if mth.Integration == nil { + return cerrors.New(cerrors.NotFound, msgIntegrationNotFound) + } + + mr, ok := mth.MethodResponses[statusCode] + if !ok { + return cerrors.New(cerrors.NotFound, msgResponseNotFound) + } + + ir := &driver.IntegrationResponse{ + StatusCode: statusCode, + SelectionPattern: in.SelectionPattern, + ResponseParameters: copyStrMap(in.ResponseParameters), + ResponseTemplates: copyStrMap(in.ResponseTemplates), + ContentHandling: in.ContentHandling, + } + + if err := validateIntegrationResponse(ir, mr); err != nil { + return err + } + + if mth.Integration.IntegrationResponses == nil { + mth.Integration.IntegrationResponses = map[string]*driver.IntegrationResponse{} + } + + mth.Integration.IntegrationResponses[statusCode] = ir + out = copyIntegrationResponse(ir) + + return nil + }) + if err != nil { + return nil, err + } + + return &out, nil +} + +// GetIntegrationResponse returns an integration response. +func (m *Mock) GetIntegrationResponse( + _ context.Context, restAPIID, resourceID, httpMethod, statusCode string, +) (*driver.IntegrationResponse, error) { + mth, err := m.lookupMethod(restAPIID, resourceID, httpMethod) + if err != nil { + return nil, err + } + + if mth.Integration == nil { + return nil, cerrors.New(cerrors.NotFound, msgIntegrationNotFound) + } + + ir, ok := mth.Integration.IntegrationResponses[statusCode] + if !ok { + return nil, cerrors.New(cerrors.NotFound, msgResponseNotFound) + } + + return ir, nil +} + +// UpdateIntegrationResponse patches an integration response. +func (m *Mock) UpdateIntegrationResponse( + _ context.Context, restAPIID, resourceID, httpMethod, statusCode string, ops []driver.PatchOperation, +) (*driver.IntegrationResponse, error) { + var out driver.IntegrationResponse + + err := m.withMethod(restAPIID, resourceID, httpMethod, func(mth *driver.Method) error { + if mth.Integration == nil { + return cerrors.New(cerrors.NotFound, msgIntegrationNotFound) + } + + ir, ok := mth.Integration.IntegrationResponses[statusCode] + if !ok { + return cerrors.New(cerrors.NotFound, msgResponseNotFound) + } + + next := copyIntegrationResponse(ir) + for _, op := range ops { + applyIntegrationResponsePatch(&next, op) + } + + if err := validateContentHandling(next.ContentHandling); err != nil { + return err + } + + if err := validateIntegrationResponse(&next, mth.MethodResponses[statusCode]); err != nil { + return err + } + + *ir = next + out = copyIntegrationResponse(ir) + + return nil + }) + if err != nil { + return nil, err + } + + return &out, nil +} + +func applyIntegrationResponsePatch(ir *driver.IntegrationResponse, op driver.PatchOperation) { + switch op.Path { + case "/selectionPattern": + ir.SelectionPattern = patchRef(op) + case pathContentHandling: + ir.ContentHandling = patchRef(op) + default: + applyMapPatch(op, "/responseTemplates/", func(k, v string, remove bool) { + ir.ResponseTemplates = patchStrMap(ir.ResponseTemplates, k, v, remove) + }) + applyMapPatch(op, "/responseParameters/", func(k, v string, remove bool) { + ir.ResponseParameters = patchStrMap(ir.ResponseParameters, k, v, remove) + }) + } +} + +// DeleteIntegrationResponse removes an integration response. +func (m *Mock) DeleteIntegrationResponse(_ context.Context, restAPIID, resourceID, httpMethod, statusCode string) error { + return m.withMethod(restAPIID, resourceID, httpMethod, func(mth *driver.Method) error { + if mth.Integration == nil { + return cerrors.New(cerrors.NotFound, msgIntegrationNotFound) + } + + if _, ok := mth.Integration.IntegrationResponses[statusCode]; !ok { + return cerrors.New(cerrors.NotFound, msgResponseNotFound) + } + + delete(mth.Integration.IntegrationResponses, statusCode) + + return nil + }) +} + +// withMethod runs fn on the live method under the API's write lock. +func (m *Mock) withMethod(restAPIID, resourceID, httpMethod string, fn func(*driver.Method) error) error { + ad, err := m.getAPI(restAPIID) + if err != nil { + return err + } + + ad.mu.Lock() + defer ad.mu.Unlock() + + res, ok := ad.resources[resourceID] + if !ok { + return cerrors.New(cerrors.NotFound, msgResourceNotFound) + } + + mth, ok := res.Methods[normalizeMethod(httpMethod)] + if !ok { + return cerrors.New(cerrors.NotFound, msgMethodNotFound) + } + + return fn(mth) +} + +// applyMapPatch calls set for an op whose path is prefix + an escaped map key. +// remove is true for a remove op. +func applyMapPatch(op driver.PatchOperation, prefix string, set func(key, value string, remove bool)) { + if len(op.Path) <= len(prefix) || op.Path[:len(prefix)] != prefix { + return + } + + set(unescapePointer(op.Path[len(prefix):]), op.Value, op.Op == opRemove) +} + +func patchStrMap(m map[string]string, k, v string, remove bool) map[string]string { + if remove { + delete(m, k) + + return m + } + + if m == nil { + m = map[string]string{} + } + + m[k] = v + + return m +} + +func patchBoolMap(m map[string]bool, k, v string, remove bool) map[string]bool { + if remove { + delete(m, k) + + return m + } + + if m == nil { + m = map[string]bool{} + } + + m[k] = parseBool(v) + + return m +} + +func copyBoolMap(in map[string]bool) map[string]bool { + if in == nil { + return nil + } + + out := make(map[string]bool, len(in)) + for k, v := range in { + out[k] = v + } + + return out +} + +func copyMethodResponse(mr *driver.MethodResponse) driver.MethodResponse { + return driver.MethodResponse{ + StatusCode: mr.StatusCode, + ResponseParameters: copyBoolMap(mr.ResponseParameters), + ResponseModels: copyStrMap(mr.ResponseModels), + } +} + +func copyIntegrationResponse(ir *driver.IntegrationResponse) driver.IntegrationResponse { + out := *ir + out.ResponseParameters = copyStrMap(ir.ResponseParameters) + out.ResponseTemplates = copyStrMap(ir.ResponseTemplates) + + return out +} + +// copyIntegration deep-copies an integration and its responses. +func copyIntegration(ig *driver.Integration) driver.Integration { + out := *ig + out.RequestParameters = copyStrMap(ig.RequestParameters) + out.RequestTemplates = copyStrMap(ig.RequestTemplates) + out.CacheKeyParameters = append([]string(nil), ig.CacheKeyParameters...) + + if ig.IntegrationResponses != nil { + out.IntegrationResponses = make(map[string]*driver.IntegrationResponse, len(ig.IntegrationResponses)) + + for code, ir := range ig.IntegrationResponses { + cp := copyIntegrationResponse(ir) + out.IntegrationResponses[code] = &cp + } + } + + return out +} + +// copyMethod deep-copies a method, its responses and its integration. +func copyMethod(mth *driver.Method) driver.Method { + out := *mth + out.RequestParameters = copyBoolMap(mth.RequestParameters) + out.RequestModels = copyStrMap(mth.RequestModels) + + if mth.MethodResponses != nil { + out.MethodResponses = make(map[string]*driver.MethodResponse, len(mth.MethodResponses)) + + for code, mr := range mth.MethodResponses { + cp := copyMethodResponse(mr) + out.MethodResponses[code] = &cp + } + } + + if mth.Integration != nil { + ig := copyIntegration(mth.Integration) + out.Integration = &ig + } + + return out +} + +// sortedResponseCodes returns an integration's response status codes in +// ascending order, so selection is deterministic. +func sortedResponseCodes(irs map[string]*driver.IntegrationResponse) []string { + codes := make([]string, 0, len(irs)) + for c := range irs { + codes = append(codes, c) + } + + sort.Strings(codes) + + return codes +} diff --git a/providers/aws/sfn/asl/jsonpath.go b/providers/aws/sfn/asl/jsonpath.go index e224133b4..56b6c7fc8 100644 --- a/providers/aws/sfn/asl/jsonpath.go +++ b/providers/aws/sfn/asl/jsonpath.go @@ -1,140 +1,15 @@ package asl -import ( - "strconv" - "strings" -) +import "github.com/stackshy/cloudemu/v2/internal/jsonpath" // evalPath evaluates a JSONPath reference against root, returning the selected -// value and whether it was present. Only the reference/selection subset is -// supported: "$", "$.a.b", "$[0]", "$.a.b[2]". Filters, wildcards and recursive -// descent are rejected with an error so an unsupported path fails loudly rather -// than returning a wrong silent result. +// value and whether it was present. The supported subset lives in +// internal/jsonpath; its errors surface as ASL definition errors. func evalPath(path string, root any) (value any, present bool, err error) { - if path == "" || path[0] != '$' { - return nil, false, aslErrf("invalid JSONPath %q: must start with '$'", path) - } - - if rerr := rejectUnsupportedPath(path); rerr != nil { - return nil, false, rerr - } - - if path == "$" { - return root, true, nil - } - - toks, terr := tokenizePath(path[1:]) - if terr != nil { - return nil, false, terr - } - - cur := root - - for _, t := range toks { - next, ok := t.apply(cur) - if !ok { - return nil, false, nil - } - - cur = next - } - - return cur, true, nil -} - -func rejectUnsupportedPath(path string) error { - if strings.ContainsAny(path, "*?@") || strings.Contains(path, "..") { - return aslErrf("JSONPath %q uses unsupported syntax (filters/wildcards/recursive descent)", path) - } - - return nil -} - -// pathToken is one selection step: a map field or a slice index. -type pathToken struct { - field string - index int - isIndex bool -} - -func (t pathToken) apply(cur any) (any, bool) { - if t.isIndex { - arr, ok := cur.([]any) - if !ok || t.index < 0 || t.index >= len(arr) { - return nil, false - } - - return arr[t.index], true - } - - m, ok := cur.(map[string]any) - if !ok { - return nil, false - } - - v, ok := m[t.field] - - return v, ok -} - -// tokenizePath splits the portion of a path after the leading '$' into tokens. -func tokenizePath(s string) ([]pathToken, error) { - var toks []pathToken - - for s != "" { - switch s[0] { - case '.': - field, rest := scanField(s[1:]) - if field == "" { - return nil, aslErrf("empty field name in JSONPath") - } - - toks = append(toks, pathToken{field: field}) - s = rest - case '[': - tok, rest, err := scanBracket(s) - if err != nil { - return nil, err - } - - toks = append(toks, tok) - s = rest - default: - return nil, aslErrf("unexpected character %q in JSONPath", s[0]) - } - } - - return toks, nil -} - -// scanField reads a dotted field name up to the next '.' or '['. -func scanField(s string) (field, rest string) { - i := strings.IndexAny(s, ".[") - if i < 0 { - return s, "" - } - - return s[:i], s[i:] -} - -// scanBracket reads a "[...]" selector: a numeric index, or a quoted field name. -func scanBracket(s string) (pathToken, string, error) { - end := strings.IndexByte(s, ']') - if end < 0 { - return pathToken{}, "", aslErrf("unterminated '[' in JSONPath") - } - - inner := s[1:end] - rest := s[end+1:] - - if len(inner) >= 2 && (inner[0] == '\'' || inner[0] == '"') { - return pathToken{field: inner[1 : len(inner)-1]}, rest, nil - } - - idx, err := strconv.Atoi(inner) + v, ok, err := jsonpath.Eval(path, root) if err != nil { - return pathToken{}, "", aslErrf("invalid array index %q in JSONPath", inner) + return nil, false, aslErrf("%s", err.Error()) } - return pathToken{index: idx, isIndex: true}, rest, nil + return v, ok, nil } diff --git a/server/aws/apigateway/handler.go b/server/aws/apigateway/handler.go index d8021cd58..51883e985 100644 --- a/server/aws/apigateway/handler.go +++ b/server/aws/apigateway/handler.go @@ -54,6 +54,8 @@ const ( segsDocItem = 4 // {id}/documentation/{parts|versions}/{item} segsMethod = 5 // {id}/resources/{rid}/methods/{httpMethod} segsIntegration = 6 // {id}/resources/{rid}/methods/{httpMethod}/integration + segsMethodResp = 7 // {id}/resources/{rid}/methods/{httpMethod}/responses/{code} + segsIntegResp = 8 // {id}/resources/{rid}/methods/{httpMethod}/integration/responses/{code} ) // Handler serves API Gateway requests against a driver. @@ -137,6 +139,10 @@ func (h *Handler) serveControlPlane(w http.ResponseWriter, r *http.Request) { h.serveMethod(w, r, segs) case segsIntegration: h.serveIntegration(w, r, segs) + case segsMethodResp: + h.serveMethodResponse(w, r, segs) + case segsIntegResp: + h.serveIntegrationResponse(w, r, segs) default: writeError(w, http.StatusNotFound, "NotFoundException", "unsupported API Gateway path") } @@ -401,6 +407,8 @@ func (h *Handler) serveMethod(w http.ResponseWriter, r *http.Request, segs []str mth, err := h.ag.PutMethod(r.Context(), id, resourceID, httpMethod, driver.PutMethodInput{ AuthorizationType: req.AuthorizationType, APIKeyRequired: req.APIKeyRequired, + OperationName: req.OperationName, RequestParameters: req.RequestParameters, + RequestModels: req.RequestModels, }) if err != nil { writeErr(w, err) @@ -432,7 +440,7 @@ func (h *Handler) serveMethod(w http.ResponseWriter, r *http.Request, segs []str // /restapis/{id}/resources/{rid}/methods/{httpMethod}/integration: // PUT=PutIntegration, GET=GetIntegration. func (h *Handler) serveIntegration(w http.ResponseWriter, r *http.Request, segs []string) { - if segs[5] != "integration" { + if segs[5] != segIntegration { writeError(w, http.StatusNotFound, "NotFoundException", "unsupported API Gateway path") return } @@ -458,7 +466,10 @@ func (h *Handler) serveIntegration(w http.ResponseWriter, r *http.Request, segs ig, err := h.ag.PutIntegration(r.Context(), id, resourceID, httpMethod, driver.PutIntegrationInput{ Type: req.Type, IntegrationHTTPMethod: req.IntegrationHTTPMethod, URI: req.URI, PassthroughBehavior: req.PassthroughBehavior, - TimeoutInMillis: req.TimeoutInMillis, + TimeoutInMillis: req.TimeoutInMillis, Credentials: req.Credentials, + RequestParameters: req.RequestParameters, RequestTemplates: req.RequestTemplates, + ContentHandling: req.ContentHandling, CacheNamespace: req.CacheNamespace, + CacheKeyParameters: req.CacheKeyParameters, }) if err != nil { writeErr(w, err) diff --git a/server/aws/apigateway/mock_integration_e2e_test.go b/server/aws/apigateway/mock_integration_e2e_test.go new file mode 100644 index 000000000..7cc9cbe1e --- /dev/null +++ b/server/aws/apigateway/mock_integration_e2e_test.go @@ -0,0 +1,211 @@ +package apigateway_test + +import ( + "context" + "io" + "net/http" + "strings" + "testing" +) + +// invoke sends a data-plane request and returns status, headers and body. +func invoke(t *testing.T, method, url, contentType, accept, body string) (int, http.Header, string) { + t.Helper() + + req, _ := http.NewRequestWithContext(context.Background(), method, url, strings.NewReader(body)) + if contentType != "" { + req.Header.Set("Content-Type", contentType) + } + + if accept != "" { + req.Header.Set("Accept", accept) + } + + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("%s %s: %v", method, url, err) + } + + defer resp.Body.Close() + + raw, _ := io.ReadAll(resp.Body) + + return resp.StatusCode, resp.Header, string(raw) +} + +// buildMockAPI creates /pets with a MOCK GET whose request template picks the +// status from ?code=, three method responses and two integration responses, +// plus a MOCK POST that never passes unmapped content through. It deploys to +// stage "test" with a stage variable and returns the API id and the +// method base URL. +func buildMockAPI(t *testing.T, base string) (apiID, resourceID string) { + t.Helper() + + api := doJSON(t, http.MethodPost, base+"/restapis", `{"name":"mock"}`) + apiID, _ = api["id"].(string) + rootID, _ := api["rootResourceId"].(string) + + res := doJSON(t, http.MethodPost, base+"/restapis/"+apiID+"/resources/"+rootID, `{"pathPart":"pets"}`) + resourceID, _ = res["id"].(string) + + m := base + "/restapis/" + apiID + "/resources/" + resourceID + "/methods/GET" + doJSON(t, http.MethodPut, m, `{"authorizationType":"NONE","requestParameters":{"method.request.querystring.code":false}}`) + doJSON(t, http.MethodPut, m+"/integration", `{"type":"MOCK","requestTemplates":{"application/json":`+ + `"{\"statusCode\": #if($input.params('code') != \"\")$input.params('code')#{else}200#end}"}}`) + doJSON(t, http.MethodPut, m+"/responses/200", `{"responseParameters":{"method.response.header.X-Custom":false}}`) + doJSON(t, http.MethodPut, m+"/responses/404", `{}`) + doJSON(t, http.MethodPut, m+"/responses/500", `{}`) + doJSON(t, http.MethodPut, m+"/integration/responses/200", `{"selectionPattern":"",`+ + `"responseParameters":{"method.response.header.X-Custom":"'yes'"},`+ + `"responseTemplates":{"application/json":`+ + `"{\"name\":\"$input.params('name')\",\"stage\":\"$context.stage\",\"v\":\"$stageVariables.v\"}"}}`) + doJSON(t, http.MethodPut, m+"/integration/responses/404", `{"selectionPattern":"4\\d\\d",`+ + `"responseTemplates":{"application/json":"{\"error\":\"not found\"}"}}`) + + p := base + "/restapis/" + apiID + "/resources/" + resourceID + "/methods/POST" + doJSON(t, http.MethodPut, p, `{"authorizationType":"NONE"}`) + doJSON(t, http.MethodPut, p+"/integration", `{"type":"MOCK","passthroughBehavior":"NEVER",`+ + `"requestTemplates":{"application/json":"{\"statusCode\": 200}"}}`) + doJSON(t, http.MethodPut, p+"/responses/200", `{}`) + doJSON(t, http.MethodPut, p+"/integration/responses/200", `{"responseTemplates":{"application/json":"$input.json('$')"}}`) + + doJSON(t, http.MethodPost, base+"/restapis/"+apiID+"/deployments", `{"stageName":"test","variables":{"v":"one"}}`) + + return apiID, resourceID +} + +func TestMockIntegrationInvoke(t *testing.T) { + srv := newE2E(t) + apiID, _ := buildMockAPI(t, srv.URL) + stage := srv.URL + "/restapis/" + apiID + "/test/_user_request_/pets" + + status, hdr, body := invoke(t, http.MethodGet, stage+"?name=rex", "", "", "") + if status != http.StatusOK || body != `{"name":"rex","stage":"test","v":"one"}` { + t.Fatalf("GET 200 = %d %s", status, body) + } + + if hdr.Get("X-Custom") != "yes" || hdr.Get("Content-Type") != "application/json" || hdr.Get("X-Amzn-Requestid") == "" { + t.Fatalf("GET 200 headers = %v", hdr) + } + + status, _, body = invoke(t, http.MethodGet, stage+"?code=404", "", "", "") + if status != http.StatusNotFound || body != `{"error":"not found"}` { + t.Fatalf("GET 404 = %d %s", status, body) + } + + // 500 has a method response but no integration response matches it and + // 200's default does: the default wins and maps to 200. + status, _, _ = invoke(t, http.MethodGet, stage+"?code=500", "", "", "") + if status != http.StatusOK { + t.Fatalf("GET 500 falls to default = %d", status) + } + + status, hdr, body = invoke(t, http.MethodPost, stage, "application/json", "", `{"a":1,"b":[true]}`) + if status != http.StatusOK || body != "{}" { + t.Fatalf("POST json = %d %s", status, body) + } + + status, hdr, body = invoke(t, http.MethodPost, stage, "text/plain", "", "hello") + if status != http.StatusUnsupportedMediaType || body != `{"message": "Unsupported Media Type"}` || + hdr.Get("X-Amzn-Errortype") != "UnsupportedMediaTypeException" { + t.Fatalf("POST text/plain = %d %s %v", status, body, hdr) + } +} + +func TestMockIntegrationNoMatchingResponse(t *testing.T) { + srv := newE2E(t) + apiID, resourceID := buildMockAPI(t, srv.URL) + m := srv.URL + "/restapis/" + apiID + "/resources/" + resourceID + "/methods/GET" + + // Drop the default response: a 200 from the mock now matches nothing. + deleteOK(t, m+"/integration/responses/200") + doJSON(t, http.MethodPost, srv.URL+"/restapis/"+apiID+"/deployments", `{"stageName":"test"}`) + + status, hdr, body := invoke(t, http.MethodGet, srv.URL+"/restapis/"+apiID+"/test/_user_request_/pets", "", "", "") + if status != http.StatusInternalServerError || body != `{"message": "Internal server error"}` || + hdr.Get("X-Amzn-Errortype") != "InternalServerErrorException" { + t.Fatalf("no match = %d %s %v", status, body, hdr) + } +} + +func TestMockIntegrationUpdateAndRedeploy(t *testing.T) { + srv := newE2E(t) + apiID, resourceID := buildMockAPI(t, srv.URL) + m := srv.URL + "/restapis/" + apiID + "/resources/" + resourceID + "/methods/GET" + + ir := doJSON(t, http.MethodPatch, m+"/integration/responses/200", `{"patchOperations":[`+ + `{"op":"replace","path":"/responseTemplates/application~1json","value":"{\"changed\":true}"}]}`) + if tm, _ := ir["responseTemplates"].(map[string]any); tm["application/json"] != `{"changed":true}` { + t.Fatalf("UpdateIntegrationResponse = %v", ir) + } + + stage := srv.URL + "/restapis/" + apiID + "/test/_user_request_/pets" + if _, _, body := invoke(t, http.MethodGet, stage, "", "", ""); body == `{"changed":true}` { + t.Fatal("live edit visible before redeploy") + } + + doJSON(t, http.MethodPost, srv.URL+"/restapis/"+apiID+"/deployments", `{"stageName":"test"}`) + + if _, _, body := invoke(t, http.MethodGet, stage, "", "", ""); body != `{"changed":true}` { + t.Fatalf("after redeploy body = %s", body) + } + + mth := doJSON(t, http.MethodGet, m, "") + mrs, _ := mth["methodResponses"].(map[string]any) + ig, _ := mth["methodIntegration"].(map[string]any) + irs, _ := ig["integrationResponses"].(map[string]any) + + if len(mrs) != 3 || len(irs) != 2 || ig["cacheNamespace"] != resourceID { + t.Fatalf("GetMethod = %v", mth) + } + + mr := doJSON(t, http.MethodPatch, m+"/responses/404", `{"patchOperations":[`+ + `{"op":"add","path":"/responseModels/application~1json","value":"Empty"}]}`) + if models, _ := mr["responseModels"].(map[string]any); models["application/json"] != "Empty" { + t.Fatalf("UpdateMethodResponse = %v", mr) + } + + ig = doJSON(t, http.MethodPatch, m+"/integration", `{"patchOperations":[`+ + `{"op":"add","path":"/requestTemplates/text~1plain","value":"{}"},`+ + `{"op":"replace","path":"/passthroughBehavior","value":"WHEN_NO_TEMPLATES"}]}`) + if tm, _ := ig["requestTemplates"].(map[string]any); len(tm) != 2 || ig["passthroughBehavior"] != "WHEN_NO_TEMPLATES" { + t.Fatalf("UpdateIntegration = %v", ig) + } +} + +func TestMethodAndIntegrationResponseErrors(t *testing.T) { + srv := newE2E(t) + apiID, resourceID := buildMockAPI(t, srv.URL) + m := srv.URL + "/restapis/" + apiID + "/resources/" + resourceID + "/methods/GET" + + const ( + badRequest = "BadRequestException" + notFound = "NotFoundException" + ) + + assertWireError(t, http.MethodPut, m+"/responses/200", `{}`, http.StatusConflict, "ConflictException", + "Response already exists for this resource") + assertWireError(t, http.MethodPut, m+"/responses/99", `{}`, http.StatusBadRequest, badRequest, "Invalid status code specified") + assertWireError(t, http.MethodPut, m+"/integration/responses/201", `{}`, http.StatusNotFound, notFound, + "Invalid Response status code specified") + assertWireError(t, http.MethodGet, m+"/responses/418", "", http.StatusNotFound, notFound, + "Invalid Response status code specified") + assertWireError(t, http.MethodPut, m+"/integration/responses/404", + `{"responseParameters":{"method.response.header.X-Nope":"'x'"}}`, http.StatusBadRequest, badRequest, + "Invalid mapping expression specified: Validation Result: warnings : [], errors : "+ + "[Invalid mapping expression parameter specified: method.response.header.X-Nope]") + assertWireError(t, http.MethodPut, m+"/integration/responses/404", `{"selectionPattern":"("}`, + http.StatusBadRequest, badRequest, "Invalid selection pattern specified") + assertWireError(t, http.MethodPatch, m+"/integration", `{"patchOperations":[`+ + `{"op":"add","path":"/requestParameters/integration.request.header.X","value":"method.request.header.Undeclared"}]}`, + http.StatusBadRequest, badRequest, "Invalid mapping expression specified: Validation Result: warnings : [], errors : "+ + "[Invalid mapping expression parameter specified: method.request.header.Undeclared]") + assertWireError(t, http.MethodPatch, m+"/integration", `{"patchOperations":[`+ + `{"op":"replace","path":"/passthroughBehavior","value":"SOMETIMES"}]}`, + http.StatusBadRequest, badRequest, "Invalid passthrough behavior specified") + + deleteOK(t, m+"/integration/responses/404") + deleteOK(t, m+"/responses/404") + assertWireError(t, http.MethodDelete, m+"/responses/404", "", http.StatusNotFound, notFound, + "Invalid Response status code specified") +} diff --git a/server/aws/apigateway/responses.go b/server/aws/apigateway/responses.go new file mode 100644 index 000000000..324779413 --- /dev/null +++ b/server/aws/apigateway/responses.go @@ -0,0 +1,103 @@ +package apigateway + +import ( + "net/http" + + "github.com/stackshy/cloudemu/v2/services/apigateway/driver" +) + +// Path segments of the method and integration response routes. +const ( + segMethods = "methods" + segIntegration = "integration" + segResponses = "responses" +) + +// serveMethodResponse handles +// /restapis/{id}/resources/{rid}/methods/{httpMethod}/responses/{statusCode}: +// PUT=PutMethodResponse, GET=GetMethodResponse, PATCH=UpdateMethodResponse, +// DELETE=DeleteMethodResponse. +func (h *Handler) serveMethodResponse(w http.ResponseWriter, r *http.Request, segs []string) { + if segs[3] != segMethods || segs[5] != segResponses { + writeError(w, http.StatusNotFound, "NotFoundException", "unsupported API Gateway path") + return + } + + id, resourceID, httpMethod, code := segs[0], segs[2], segs[4], segs[6] + ctx := r.Context() + + if r.Method == http.MethodPut { + var req putMethodResponseRequest + if !decodeJSON(w, r, &req) { + return + } + + mr, err := h.ag.PutMethodResponse(ctx, id, resourceID, httpMethod, code, driver.PutMethodResponseInput{ + ResponseParameters: req.ResponseParameters, ResponseModels: req.ResponseModels, + }) + if err != nil { + writeErr(w, err) + return + } + + writeJSON(w, http.StatusCreated, toMethodResponseObject(mr)) + + return + } + + serveItem(w, r, + func(ops []driver.PatchOperation) (*driver.MethodResponse, error) { + return h.ag.UpdateMethodResponse(ctx, id, resourceID, httpMethod, code, ops) + }, + func() (*driver.MethodResponse, error) { + return h.ag.GetMethodResponse(ctx, id, resourceID, httpMethod, code) + }, + func() error { return h.ag.DeleteMethodResponse(ctx, id, resourceID, httpMethod, code) }, + toMethodResponseObject, + ) +} + +// serveIntegrationResponse handles +// /restapis/{id}/resources/{rid}/methods/{httpMethod}/integration/responses/{statusCode}: +// PUT=PutIntegrationResponse, GET=GetIntegrationResponse, +// PATCH=UpdateIntegrationResponse, DELETE=DeleteIntegrationResponse. +func (h *Handler) serveIntegrationResponse(w http.ResponseWriter, r *http.Request, segs []string) { + if segs[3] != segMethods || segs[5] != segIntegration || segs[6] != segResponses { + writeError(w, http.StatusNotFound, "NotFoundException", "unsupported API Gateway path") + return + } + + id, resourceID, httpMethod, code := segs[0], segs[2], segs[4], segs[7] + ctx := r.Context() + + if r.Method == http.MethodPut { + var req putIntegrationResponseRequest + if !decodeJSON(w, r, &req) { + return + } + + ir, err := h.ag.PutIntegrationResponse(ctx, id, resourceID, httpMethod, code, driver.PutIntegrationResponseInput{ + SelectionPattern: req.SelectionPattern, ResponseParameters: req.ResponseParameters, + ResponseTemplates: req.ResponseTemplates, ContentHandling: req.ContentHandling, + }) + if err != nil { + writeErr(w, err) + return + } + + writeJSON(w, http.StatusCreated, toIntegrationResponseObject(ir)) + + return + } + + serveItem(w, r, + func(ops []driver.PatchOperation) (*driver.IntegrationResponse, error) { + return h.ag.UpdateIntegrationResponse(ctx, id, resourceID, httpMethod, code, ops) + }, + func() (*driver.IntegrationResponse, error) { + return h.ag.GetIntegrationResponse(ctx, id, resourceID, httpMethod, code) + }, + func() error { return h.ag.DeleteIntegrationResponse(ctx, id, resourceID, httpMethod, code) }, + toIntegrationResponseObject, + ) +} diff --git a/server/aws/apigateway/types.go b/server/aws/apigateway/types.go index acab026cf..aa51855c6 100644 --- a/server/aws/apigateway/types.go +++ b/server/aws/apigateway/types.go @@ -52,19 +52,28 @@ type createResourceRequest struct { // putMethodRequest is the PutMethod request body. type putMethodRequest struct { - AuthorizationType string `json:"authorizationType"` - APIKeyRequired bool `json:"apiKeyRequired"` + AuthorizationType string `json:"authorizationType"` + APIKeyRequired bool `json:"apiKeyRequired"` + OperationName string `json:"operationName"` + RequestParameters map[string]bool `json:"requestParameters"` + RequestModels map[string]string `json:"requestModels"` } // putIntegrationRequest is the PutIntegration request body. The integration's // backend method travels as "httpMethod" on the wire (the model's locationName // for integrationHttpMethod). type putIntegrationRequest struct { - Type string `json:"type"` - IntegrationHTTPMethod string `json:"httpMethod"` - URI string `json:"uri"` - PassthroughBehavior string `json:"passthroughBehavior"` - TimeoutInMillis int `json:"timeoutInMillis"` + Type string `json:"type"` + IntegrationHTTPMethod string `json:"httpMethod"` + URI string `json:"uri"` + PassthroughBehavior string `json:"passthroughBehavior"` + TimeoutInMillis int `json:"timeoutInMillis"` + Credentials string `json:"credentials"` + RequestParameters map[string]string `json:"requestParameters"` + RequestTemplates map[string]string `json:"requestTemplates"` + ContentHandling string `json:"contentHandling"` + CacheNamespace string `json:"cacheNamespace"` + CacheKeyParameters []string `json:"cacheKeyParameters"` } // createDeploymentRequest is the CreateDeployment request body. @@ -128,19 +137,60 @@ type listResourcesResponse struct { // methodResponse is the Method wire object. type methodResponse struct { - HTTPMethod string `json:"httpMethod,omitempty"` - AuthorizationType string `json:"authorizationType,omitempty"` - APIKeyRequired bool `json:"apiKeyRequired"` - MethodIntegration *integrationResponse `json:"methodIntegration,omitempty"` + HTTPMethod string `json:"httpMethod,omitempty"` + AuthorizationType string `json:"authorizationType,omitempty"` + APIKeyRequired bool `json:"apiKeyRequired"` + OperationName string `json:"operationName,omitempty"` + RequestParameters map[string]bool `json:"requestParameters,omitempty"` + RequestModels map[string]string `json:"requestModels,omitempty"` + MethodResponses map[string]methodResponseObject `json:"methodResponses,omitempty"` + MethodIntegration *integrationResponse `json:"methodIntegration,omitempty"` } // integrationResponse is the Integration wire object. type integrationResponse struct { - Type string `json:"type"` - HTTPMethod string `json:"httpMethod,omitempty"` - URI string `json:"uri,omitempty"` - PassthroughBehavior string `json:"passthroughBehavior,omitempty"` - TimeoutInMillis int `json:"timeoutInMillis,omitempty"` + Type string `json:"type"` + HTTPMethod string `json:"httpMethod,omitempty"` + URI string `json:"uri,omitempty"` + PassthroughBehavior string `json:"passthroughBehavior,omitempty"` + TimeoutInMillis int `json:"timeoutInMillis,omitempty"` + Credentials string `json:"credentials,omitempty"` + RequestParameters map[string]string `json:"requestParameters,omitempty"` + RequestTemplates map[string]string `json:"requestTemplates,omitempty"` + ContentHandling string `json:"contentHandling,omitempty"` + CacheNamespace string `json:"cacheNamespace,omitempty"` + CacheKeyParameters []string `json:"cacheKeyParameters"` + IntegrationResponses map[string]integrationResponseObject `json:"integrationResponses,omitempty"` +} + +// methodResponseObject is the MethodResponse wire object. +type methodResponseObject struct { + StatusCode string `json:"statusCode"` + ResponseParameters map[string]bool `json:"responseParameters,omitempty"` + ResponseModels map[string]string `json:"responseModels,omitempty"` +} + +// integrationResponseObject is the IntegrationResponse wire object. +type integrationResponseObject struct { + StatusCode string `json:"statusCode"` + SelectionPattern string `json:"selectionPattern,omitempty"` + ResponseParameters map[string]string `json:"responseParameters,omitempty"` + ResponseTemplates map[string]string `json:"responseTemplates,omitempty"` + ContentHandling string `json:"contentHandling,omitempty"` +} + +// putMethodResponseRequest is the PutMethodResponse request body. +type putMethodResponseRequest struct { + ResponseParameters map[string]bool `json:"responseParameters"` + ResponseModels map[string]string `json:"responseModels"` +} + +// putIntegrationResponseRequest is the PutIntegrationResponse request body. +type putIntegrationResponseRequest struct { + SelectionPattern string `json:"selectionPattern"` + ResponseParameters map[string]string `json:"responseParameters"` + ResponseTemplates map[string]string `json:"responseTemplates"` + ContentHandling string `json:"contentHandling"` } // deploymentResponse is the Deployment wire object. APISummary is only sent @@ -229,7 +279,15 @@ func renderResource(r *driver.Resource, embedMethods bool) resourceResponse { func toMethodResponse(mth *driver.Method) methodResponse { resp := methodResponse{ HTTPMethod: mth.HTTPMethod, AuthorizationType: mth.AuthorizationType, - APIKeyRequired: mth.APIKeyRequired, + APIKeyRequired: mth.APIKeyRequired, OperationName: mth.OperationName, + RequestParameters: mth.RequestParameters, RequestModels: mth.RequestModels, + } + + if len(mth.MethodResponses) > 0 { + resp.MethodResponses = make(map[string]methodResponseObject, len(mth.MethodResponses)) + for code, mr := range mth.MethodResponses { + resp.MethodResponses[code] = toMethodResponseObject(mr) + } } if mth.Integration != nil { @@ -241,10 +299,40 @@ func toMethodResponse(mth *driver.Method) methodResponse { } func toIntegrationResponse(ig *driver.Integration) integrationResponse { - return integrationResponse{ + resp := integrationResponse{ Type: ig.Type, HTTPMethod: ig.IntegrationHTTPMethod, URI: ig.URI, PassthroughBehavior: ig.PassthroughBehavior, - TimeoutInMillis: ig.TimeoutInMillis, + TimeoutInMillis: ig.TimeoutInMillis, Credentials: ig.Credentials, + RequestParameters: ig.RequestParameters, RequestTemplates: ig.RequestTemplates, + ContentHandling: ig.ContentHandling, CacheNamespace: ig.CacheNamespace, + CacheKeyParameters: ig.CacheKeyParameters, + } + + if resp.CacheKeyParameters == nil { + resp.CacheKeyParameters = []string{} + } + + if len(ig.IntegrationResponses) > 0 { + resp.IntegrationResponses = make(map[string]integrationResponseObject, len(ig.IntegrationResponses)) + for code, ir := range ig.IntegrationResponses { + resp.IntegrationResponses[code] = toIntegrationResponseObject(ir) + } + } + + return resp +} + +func toMethodResponseObject(mr *driver.MethodResponse) methodResponseObject { + return methodResponseObject{ + StatusCode: mr.StatusCode, ResponseParameters: mr.ResponseParameters, ResponseModels: mr.ResponseModels, + } +} + +func toIntegrationResponseObject(ir *driver.IntegrationResponse) integrationResponseObject { + return integrationResponseObject{ + StatusCode: ir.StatusCode, SelectionPattern: ir.SelectionPattern, + ResponseParameters: ir.ResponseParameters, ResponseTemplates: ir.ResponseTemplates, + ContentHandling: ir.ContentHandling, } } diff --git a/services/apigateway/driver/driver.go b/services/apigateway/driver/driver.go index 47c9429a8..e44fbd557 100644 --- a/services/apigateway/driver/driver.go +++ b/services/apigateway/driver/driver.go @@ -72,23 +72,70 @@ type Resource struct { } // Method is an HTTP method configured on a Resource, optionally wired to an -// Integration. +// Integration. RequestParameters maps a method.request.{location}.{name} +// expression to whether it is required; RequestModels maps a content type to +// a model name. MethodResponses is keyed by status code. type Method struct { HTTPMethod string AuthorizationType string APIKeyRequired bool + OperationName string + RequestParameters map[string]bool + RequestModels map[string]string + MethodResponses map[string]*MethodResponse Integration *Integration } +// MethodResponse declares a status code a method can return, the response +// headers it may carry (method.response.header.{name} -> required) and the +// models of its body per content type. +type MethodResponse struct { + StatusCode string + ResponseParameters map[string]bool + ResponseModels map[string]string +} + +// IntegrationResponse maps a backend response to a method response. +// SelectionPattern is a regular expression matched against the backend status +// code (HTTP and MOCK) or Lambda error message; the response with an empty +// pattern is the default. ResponseParameters maps +// method.response.header.{name} to a source expression and ResponseTemplates +// maps a content type to a VTL template. +type IntegrationResponse struct { + StatusCode string + SelectionPattern string + ResponseParameters map[string]string + ResponseTemplates map[string]string + ContentHandling string +} + +// Integration passthrough behaviors. +const ( + PassthroughWhenNoMatch = "WHEN_NO_MATCH" + PassthroughWhenNoTemplates = "WHEN_NO_TEMPLATES" + PassthroughNever = "NEVER" +) + // Integration is the backend a Method forwards to. For AWS_PROXY/AWS the URI is // the Lambda invocation ARN // (arn:aws:apigateway::lambda:path/2015-03-31/functions//invocations). +// +// RequestParameters maps integration.request.{location}.{name} to a source +// expression and RequestTemplates maps a content type to a VTL template. +// IntegrationResponses is keyed by status code. type Integration struct { Type string IntegrationHTTPMethod string URI string PassthroughBehavior string TimeoutInMillis int + Credentials string + RequestParameters map[string]string + RequestTemplates map[string]string + ContentHandling string + CacheNamespace string + CacheKeyParameters []string + IntegrationResponses map[string]*IntegrationResponse } // Deployment is a point-in-time snapshot of a REST API published to a stage. @@ -143,16 +190,41 @@ type CreateRestAPIInput struct { type PutMethodInput struct { AuthorizationType string APIKeyRequired bool + OperationName string + RequestParameters map[string]bool + RequestModels map[string]string +} + +// PutMethodResponseInput carries the fields PutMethodResponse accepts. +type PutMethodResponseInput struct { + ResponseParameters map[string]bool + ResponseModels map[string]string +} + +// PutIntegrationResponseInput carries the fields PutIntegrationResponse +// accepts. +type PutIntegrationResponseInput struct { + SelectionPattern string + ResponseParameters map[string]string + ResponseTemplates map[string]string + ContentHandling string } // PutIntegrationInput carries the fields PutIntegration accepts. TimeoutInMillis // of 0 selects the AWS default (29000ms); a non-zero value is stored verbatim. +// An empty CacheNamespace defaults to the resource id. type PutIntegrationInput struct { Type string IntegrationHTTPMethod string URI string PassthroughBehavior string TimeoutInMillis int + Credentials string + RequestParameters map[string]string + RequestTemplates map[string]string + ContentHandling string + CacheNamespace string + CacheKeyParameters []string } // CreateDeploymentInput carries the fields CreateDeployment accepts. A non-empty @@ -377,6 +449,35 @@ type APIGateway interface { UpdateIntegration(ctx context.Context, restAPIID, resourceID, httpMethod string, ops []PatchOperation) (*Integration, error) DeleteIntegration(ctx context.Context, restAPIID, resourceID, httpMethod string) error + // PutMethodResponse declares a status code on a method. It fails when the + // status code is already declared. + PutMethodResponse( + ctx context.Context, restAPIID, resourceID, httpMethod, statusCode string, in PutMethodResponseInput, + ) (*MethodResponse, error) + GetMethodResponse(ctx context.Context, restAPIID, resourceID, httpMethod, statusCode string) (*MethodResponse, error) + // UpdateMethodResponse applies a patchOperations document + // (/responseParameters/{name}, /responseModels/{contentType}). + UpdateMethodResponse( + ctx context.Context, restAPIID, resourceID, httpMethod, statusCode string, ops []PatchOperation, + ) (*MethodResponse, error) + DeleteMethodResponse(ctx context.Context, restAPIID, resourceID, httpMethod, statusCode string) error + + // PutIntegrationResponse creates or replaces an integration response. The + // method must declare a method response with the same status code. + PutIntegrationResponse( + ctx context.Context, restAPIID, resourceID, httpMethod, statusCode string, in PutIntegrationResponseInput, + ) (*IntegrationResponse, error) + GetIntegrationResponse( + ctx context.Context, restAPIID, resourceID, httpMethod, statusCode string, + ) (*IntegrationResponse, error) + // UpdateIntegrationResponse applies a patchOperations document + // (/selectionPattern, /contentHandling, /responseTemplates/{contentType}, + // /responseParameters/{name}). + UpdateIntegrationResponse( + ctx context.Context, restAPIID, resourceID, httpMethod, statusCode string, ops []PatchOperation, + ) (*IntegrationResponse, error) + DeleteIntegrationResponse(ctx context.Context, restAPIID, resourceID, httpMethod, statusCode string) error + CreateDeployment(ctx context.Context, restAPIID string, in CreateDeploymentInput) (*Deployment, error) GetDeployments(ctx context.Context, restAPIID string) ([]Deployment, error) GetDeployment(ctx context.Context, restAPIID, deploymentID string) (*Deployment, error) @@ -443,7 +544,9 @@ type APIGateway interface { // InvokeRoute routes req through the tree its stage's deployment captured. // It resolves req.HTTPMethod+req.Path ({proxy+} greedy paths and {param} - // placeholders supported) and, for an AWS_PROXY/AWS Lambda integration, - // invokes the target function and returns its mapped HTTP response. + // placeholders supported). An AWS_PROXY/AWS Lambda integration invokes the + // target function and returns its mapped HTTP response; a MOCK integration + // renders its request template, selects an integration response and + // renders that response's mapping template. InvokeRoute(ctx context.Context, req *ProxyRequest) (*ProxyResponse, error) } From b79bf56b636dc8cf75aa5b405ab8e2743992feda Mon Sep 17 00:00:00 2001 From: Nitin Kumar Date: Sun, 4 Oct 2026 18:29:12 +0530 Subject: [PATCH 2/3] fix(aws-apigateway): bound VTL cycles, growth and parse cost; template cache (A2 review) --- internal/jsonpath/jsonpath.go | 284 +++++++++++++--- internal/jsonpath/jsonpath_test.go | 76 ++++- internal/vtl/ast.go | 5 +- internal/vtl/cache.go | 70 ++++ internal/vtl/cache_test.go | 55 ++++ internal/vtl/edge_test.go | 5 +- internal/vtl/eval.go | 224 ++++++++++--- internal/vtl/fuzz_test.go | 79 +++++ internal/vtl/limits.go | 112 +++++++ internal/vtl/limits_test.go | 164 ++++++++++ internal/vtl/methods.go | 172 ++++++++-- internal/vtl/parse.go | 155 ++++++--- internal/vtl/value.go | 308 +++++++++++++----- internal/vtl/vtl_test.go | 8 +- providers/aws/apigateway/apigateway.go | 9 + providers/aws/apigateway/mapping.go | 51 ++- .../aws/apigateway/mapping_limits_test.go | 102 ++++++ .../aws/apigateway/mapping_validation.go | 21 ++ providers/aws/apigateway/mock_integration.go | 5 +- 19 files changed, 1652 insertions(+), 253 deletions(-) create mode 100644 internal/vtl/cache.go create mode 100644 internal/vtl/cache_test.go create mode 100644 internal/vtl/fuzz_test.go create mode 100644 internal/vtl/limits.go create mode 100644 internal/vtl/limits_test.go create mode 100644 providers/aws/apigateway/mapping_limits_test.go diff --git a/internal/jsonpath/jsonpath.go b/internal/jsonpath/jsonpath.go index f80a56fbb..c2cfc5a80 100644 --- a/internal/jsonpath/jsonpath.go +++ b/internal/jsonpath/jsonpath.go @@ -1,8 +1,13 @@ -// Package jsonpath evaluates the reference subset of JSONPath shared by the -// Step Functions ASL interpreter and the API Gateway mapping templates: "$", -// "$.a.b", "$[0]", "$['a']" and "$.a.b[2]". Filters, wildcards and recursive -// descent are rejected with an error, so an unsupported path fails loudly -// instead of returning a wrong result. +// Package jsonpath evaluates JSONPath expressions over decoded JSON. +// +// Eval handles the definite subset the Step Functions ASL interpreter uses: +// "$", "$.a.b", "$[0]", "$['a']" and "$.a.b[2]". Filters, wildcards and +// recursive descent are rejected with an error, so an unsupported path fails +// loudly instead of returning a wrong result. +// +// EvalAll also accepts the wildcard ("[*]", ".*") and recursive descent +// ("..name", "..*") forms the API Gateway mapping templates allow, and returns +// every match. Filters stay unsupported. // // Values are plain decoded JSON (map[string]any, []any) or any type that // implements Object or Array, which lets callers keep an order-preserving @@ -11,6 +16,7 @@ package jsonpath import ( "fmt" + "sort" "strconv" "strings" ) @@ -18,13 +24,23 @@ import ( // Object is a JSON object view a path can step into by field name. type Object interface { Lookup(key string) (any, bool) + Keys() []string } // Array is a JSON array view a path can step into by index. type Array interface { Index(i int) (any, bool) + Len() int } +// maxDepth bounds how deep recursive descent walks, so a pathological document +// cannot exhaust the stack. +const maxDepth = 1000 + +// maxMatches caps the values one EvalAll step may collect, so chained +// wildcards and descents cannot multiply a document into a huge result. +const maxMatches = 1 << 20 + // Error reports a malformed or unsupported path. type Error struct { Msg string @@ -36,21 +52,17 @@ func errorf(format string, args ...any) error { return &Error{Msg: fmt.Sprintf(format, args...)} } -// Eval evaluates path against root and returns the selected value and whether -// it was present. +// Eval evaluates a definite path against root and returns the selected value +// and whether it was present. func Eval(path string, root any) (value any, present bool, err error) { - if path == "" || path[0] != '$' { - return nil, false, errorf("invalid JSONPath %q: must start with '$'", path) + if rootErr := checkRoot(path); rootErr != nil { + return nil, false, rootErr } if strings.ContainsAny(path, "*?@") || strings.Contains(path, "..") { return nil, false, errorf("JSONPath %q uses unsupported syntax (filters/wildcards/recursive descent)", path) } - if path == "$" { - return root, true, nil - } - toks, err := tokenize(path[1:]) if err != nil { return nil, false, err @@ -70,82 +82,252 @@ func Eval(path string, root any) (value any, present bool, err error) { return cur, true, nil } -// token is one selection step: a field name or an array index. +// EvalAll evaluates a path that may use wildcards or recursive descent. +// indefinite reports whether the path can match more than one value (it uses +// a wildcard or descent), in which case callers present the matches as a list. +// For a definite path values holds at most one element. +func EvalAll(path string, root any) (values []any, indefinite bool, err error) { + if rootErr := checkRoot(path); rootErr != nil { + return nil, false, rootErr + } + + if strings.ContainsAny(path, "?@") { + return nil, false, errorf("JSONPath %q uses unsupported syntax (filters)", path) + } + + toks, err := tokenize(path[1:]) + if err != nil { + return nil, false, err + } + + cur := []any{root} + + for _, t := range toks { + indefinite = indefinite || t.wildcard || t.descent + + var next []any + + for _, v := range cur { + next = t.collect(v, next) + + if len(next) > maxMatches { + return nil, true, errorf("JSONPath %q matches more than %d values", path, maxMatches) + } + } + + cur = next + } + + return cur, indefinite, nil +} + +func checkRoot(path string) error { + if path == "" || path[0] != '$' { + return errorf("invalid JSONPath %q: must start with '$'", path) + } + + return nil +} + +// token is one selection step: a field name, an array index or a wildcard, +// optionally applied at every depth (descent). type token struct { - field string - index int - isIndex bool + field string + index int + isIndex bool + wildcard bool + descent bool } +// apply selects one child of cur (definite tokens only). func (t token) apply(cur any) (any, bool) { if t.isIndex { - switch arr := cur.(type) { - case []any: - if t.index < 0 || t.index >= len(arr) { - return nil, false - } + return elem(cur, t.index) + } - return arr[t.index], true - case Array: - return arr.Index(t.index) - default: + return field(cur, t.field) +} + +// collect appends every match of t under v to out. +func (t token) collect(v any, out []any) []any { + if !t.descent { + return t.collectHere(v, out) + } + + // Walk v and all of its descendants in document order. + type frame struct { + v any + depth int + } + + stack := []frame{{v, 0}} + + for len(stack) > 0 { + f := stack[len(stack)-1] + stack = stack[:len(stack)-1] + + out = t.collectHere(f.v, out) + + if f.depth >= maxDepth { + continue + } + + kids := children(f.v) + for i := len(kids) - 1; i >= 0; i-- { + stack = append(stack, frame{kids[i], f.depth + 1}) + } + } + + return out +} + +// collectHere appends the matches of t directly under v. +func (t token) collectHere(v any, out []any) []any { + switch { + case t.wildcard: + return append(out, children(v)...) + case t.isIndex: + if got, ok := elem(v, t.index); ok { + return append(out, got) + } + default: + if got, ok := field(v, t.field); ok { + return append(out, got) + } + } + + return out +} + +func elem(cur any, i int) (any, bool) { + switch arr := cur.(type) { + case []any: + if i < 0 || i >= len(arr) { return nil, false } + + return arr[i], true + case Array: + return arr.Index(i) + default: + return nil, false } +} +func field(cur any, name string) (any, bool) { switch obj := cur.(type) { case map[string]any: - v, ok := obj[t.field] + v, ok := obj[name] return v, ok case Object: - return obj.Lookup(t.field) + return obj.Lookup(name) default: return nil, false } } +// children returns an object's values (in key order) or an array's elements. +func children(v any) []any { + switch t := v.(type) { + case []any: + return t + case Array: + out := make([]any, 0, t.Len()) + + for i := range t.Len() { + e, _ := t.Index(i) + out = append(out, e) + } + + return out + case map[string]any: + keys := make([]string, 0, len(t)) + for k := range t { + keys = append(keys, k) + } + + sort.Strings(keys) + + out := make([]any, 0, len(keys)) + for _, k := range keys { + out = append(out, t[k]) + } + + return out + case Object: + keys := t.Keys() + out := make([]any, 0, len(keys)) + + for _, k := range keys { + e, _ := t.Lookup(k) + out = append(out, e) + } + + return out + default: + return nil + } +} + // tokenize splits the part of a path after the leading '$' into tokens. func tokenize(s string) ([]token, error) { var toks []token for s != "" { - switch s[0] { - case '.': - field, rest := scanField(s[1:]) - if field == "" { - return nil, errorf("empty field name in JSONPath") - } + var ( + tok token + err error + ) - toks = append(toks, token{field: field}) - s = rest - case '[': - tok, rest, err := scanBracket(s) - if err != nil { - return nil, err - } - - toks = append(toks, tok) - s = rest + switch { + case strings.HasPrefix(s, ".."): + tok, s, err = scanStep(s[2:]) + tok.descent = true + case s[0] == '.': + tok, s, err = scanStep(s[1:]) + case s[0] == '[': + tok, s, err = scanBracket(s) default: return nil, errorf("unexpected character %q in JSONPath", s[0]) } + + if err != nil { + return nil, err + } + + toks = append(toks, tok) } return toks, nil } -// scanField reads a dotted field name up to the next '.' or '['. -func scanField(s string) (field, rest string) { +// scanStep reads the selector after a '.' or '..': a field name, '*' or a +// bracket. +func scanStep(s string) (token, string, error) { + if strings.HasPrefix(s, "[") { + return scanBracket(s) + } + i := strings.IndexAny(s, ".[") if i < 0 { - return s, "" + i = len(s) } - return s[:i], s[i:] + name := s[:i] + + switch name { + case "": + return token{}, "", errorf("empty field name in JSONPath") + case "*": + return token{wildcard: true}, s[i:], nil + default: + return token{field: name}, s[i:], nil + } } -// scanBracket reads a "[...]" selector: a numeric index or a quoted field name. +// scanBracket reads a "[...]" selector: a numeric index, '*' or a quoted +// field name. func scanBracket(s string) (token, string, error) { end := strings.IndexByte(s, ']') if end < 0 { @@ -155,6 +337,10 @@ func scanBracket(s string) (token, string, error) { inner := s[1:end] rest := s[end+1:] + if inner == "*" { + return token{wildcard: true}, rest, nil + } + if len(inner) >= 2 && (inner[0] == '\'' || inner[0] == '"') { return token{field: inner[1 : len(inner)-1]}, rest, nil } diff --git a/internal/jsonpath/jsonpath_test.go b/internal/jsonpath/jsonpath_test.go index b86177e60..4bb3cf749 100644 --- a/internal/jsonpath/jsonpath_test.go +++ b/internal/jsonpath/jsonpath_test.go @@ -1,6 +1,10 @@ package jsonpath -import "testing" +import ( + "fmt" + "sort" + "testing" +) type obj map[string]any @@ -10,6 +14,17 @@ func (o obj) Lookup(k string) (any, bool) { return v, ok } +func (o obj) Keys() []string { + keys := make([]string, 0, len(o)) + for k := range o { + keys = append(keys, k) + } + + sort.Strings(keys) + + return keys +} + type arr []any func (a arr) Index(i int) (any, bool) { @@ -20,6 +35,8 @@ func (a arr) Index(i int) (any, bool) { return a[i], true } +func (a arr) Len() int { return len(a) } + func TestEval(t *testing.T) { root := map[string]any{"a": map[string]any{"b": []any{1, 2}}, "o": obj{"x": arr{"y"}}} @@ -59,3 +76,60 @@ func TestEval(t *testing.T) { } } } + +func TestEvalAll(t *testing.T) { + root := map[string]any{ + "items": []any{map[string]any{"id": 1}, map[string]any{"id": 2, "sub": map[string]any{"id": 3}}}, + "o": obj{"b": arr{"x"}, "a": "y"}, + } + + cases := []struct { + path string + want string + indefinite bool + }{ + {"$.items[*].id", "[1 2]", true}, + {"$.items.*.id", "[1 2]", true}, + {"$..id", "[1 2 3]", true}, + {"$.o.*", "[y [x]]", true}, + {"$..[0]", "[map[id:1] x]", true}, + {"$.items[1].sub.id", "[3]", false}, + {"$.missing", "[]", false}, + {"$..nothing", "[]", true}, + } + + for _, c := range cases { + got, indefinite, err := EvalAll(c.path, root) + if err != nil || fmt.Sprint(got) != c.want && !(len(got) == 0 && c.want == "[]") || indefinite != c.indefinite { + t.Errorf("EvalAll(%q) = %v %v %v, want %s %v", c.path, got, indefinite, err, c.want, c.indefinite) + } + } + + for _, bad := range []string{"x", "$[?(@.a)]", "$..", "$.a[", "$[x]"} { + if _, _, err := EvalAll(bad, root); err == nil { + t.Errorf("EvalAll(%q) accepted", bad) + } + } + + // Descent stops at maxDepth instead of exhausting the stack. + var deep any = "leaf" + for range maxDepth + 50 { + deep = []any{deep} + } + + if _, _, err := EvalAll("$..*", deep); err != nil { + t.Fatalf("deep descent: %v", err) + } +} + +func TestEvalAllCapsMatches(t *testing.T) { + // A 900-deep chain: each descent step multiplies the matches. + var items any = 1 + for range 900 { + items = []any{items, 2} + } + + if _, _, err := EvalAll("$..*..*..*", items); err == nil { + t.Fatal("multiplying path not capped") + } +} diff --git a/internal/vtl/ast.go b/internal/vtl/ast.go index a927c9e49..d3bfffec4 100644 --- a/internal/vtl/ast.go +++ b/internal/vtl/ast.go @@ -3,8 +3,9 @@ package vtl // Template nodes. A template body is a []node rendered in order. type node any -// textNode is literal output. -type textNode struct{ text string } +// textNode is literal output, kept as the source pieces it was parsed from so +// building it stays linear. +type textNode struct{ parts []string } // refNode prints a reference. quiet is the $!x form. type refNode struct { diff --git a/internal/vtl/cache.go b/internal/vtl/cache.go new file mode 100644 index 000000000..55ef15b6a --- /dev/null +++ b/internal/vtl/cache.go @@ -0,0 +1,70 @@ +package vtl + +import ( + "container/list" + "sync" +) + +// Cache is a bounded, least-recently-used cache of parsed templates keyed by +// their source, so a template is parsed once rather than on every render. A +// parse failure is cached too. A parsed Template is read-only, so one cached +// template may render concurrently. +type Cache struct { + mu sync.Mutex + max int + order *list.List + items map[string]*list.Element +} + +type cacheEntry struct { + src string + tmpl *Template + err error +} + +// NewCache returns a cache holding at most maxEntries templates. +func NewCache(maxEntries int) *Cache { + return &Cache{max: maxEntries, order: list.New(), items: map[string]*list.Element{}} +} + +// Parse returns the parsed template for src, parsing it on a miss. +func (c *Cache) Parse(src string) (*Template, error) { + c.mu.Lock() + + if el, ok := c.items[src]; ok { + c.order.MoveToFront(el) + e, _ := el.Value.(*cacheEntry) + c.mu.Unlock() + + return e.tmpl, e.err + } + + c.mu.Unlock() + + tmpl, err := Parse(src) + + c.mu.Lock() + defer c.mu.Unlock() + + if _, ok := c.items[src]; !ok { + c.items[src] = c.order.PushFront(&cacheEntry{src: src, tmpl: tmpl, err: err}) + + for c.order.Len() > c.max { + oldest := c.order.Back() + c.order.Remove(oldest) + + e, _ := oldest.Value.(*cacheEntry) + delete(c.items, e.src) + } + } + + return tmpl, err +} + +// Len returns the number of cached templates. +func (c *Cache) Len() int { + c.mu.Lock() + defer c.mu.Unlock() + + return c.order.Len() +} diff --git a/internal/vtl/cache_test.go b/internal/vtl/cache_test.go new file mode 100644 index 000000000..01cae9c61 --- /dev/null +++ b/internal/vtl/cache_test.go @@ -0,0 +1,55 @@ +package vtl + +import ( + "context" + "sync" + "testing" +) + +func TestCache(t *testing.T) { + c := NewCache(2) + + a1, err := c.Parse("a$x") + if err != nil { + t.Fatal(err) + } + + if a2, _ := c.Parse("a$x"); a2 != a1 { + t.Fatal("cache miss on repeat") + } + + if _, err := c.Parse("#if("); err == nil { + t.Fatal("parse error not returned") + } + + if _, err := c.Parse("#if("); err == nil { + t.Fatal("cached parse error not returned") + } + + // "a$x" was used before "#if(", so adding a third entry evicts it. + if _, err := c.Parse("b"); err != nil || c.Len() != 2 { + t.Fatalf("len = %d err = %v", c.Len(), err) + } + + if a3, _ := c.Parse("a$x"); a3 == a1 { + t.Fatal("evicted entry still cached") + } + + // A cached template renders concurrently. + var wg sync.WaitGroup + + for i := range 8 { + wg.Add(1) + + go func() { + defer wg.Done() + + res, err := a1.Render(context.Background(), map[string]any{"x": int64(i)}, RenderOptions{}) + if err != nil || res.Output != "a"+Stringify(int64(i)) { + t.Errorf("render %d = %v %v", i, res, err) + } + }() + } + + wg.Wait() +} diff --git a/internal/vtl/edge_test.go b/internal/vtl/edge_test.go index 6e11826f7..5a993626b 100644 --- a/internal/vtl/edge_test.go +++ b/internal/vtl/edge_test.go @@ -104,7 +104,7 @@ func TestStringifyAndJSON(t *testing.T) { t.Fatalf("Stringify: %s", Stringify(1e21)) } - if got := ToJSON(NewList(int64(1), 2.5, true, nil, hostObj{})); got != `[1,2.5,true,null,null]` { + if got := mustJSON(t, NewList(int64(1), 2.5, true, nil, hostObj{})); got != `[1,2.5,true,null,null]` { t.Fatalf("ToJSON list = %s", got) } @@ -118,7 +118,8 @@ func TestStringifyAndJSON(t *testing.T) { t.Fatal("float parse") } - m := MapOf("a", 1) + m := NewMap() + m.Put("a", 1) if m.Remove("zz") != nil || m.Len() != 1 || len(m.Keys()) != 1 { t.Fatal("map ops") } diff --git a/internal/vtl/eval.go b/internal/vtl/eval.go index b4a68d454..9a0d76214 100644 --- a/internal/vtl/eval.go +++ b/internal/vtl/eval.go @@ -15,7 +15,7 @@ const ( // silently stops after this many iterations. MaxForeachIterations = 1000 // ctxCheckEvery is how often (in steps) the context deadline is checked. - ctxCheckEvery = 1024 + ctxCheckEvery = 64 ) // ErrStepBudget is returned when a render exceeds its step budget. @@ -37,13 +37,14 @@ type RenderOptions struct { } // Render evaluates the template with vars as its top-level references. vars is -// modified by #set. Rendering stops with ctx's error when ctx is done. +// modified by #set. Rendering stops with ctx's error when ctx is done, and +// with a limit error when the output, memory or step budget is exhausted. func (t *Template) Render(ctx context.Context, vars map[string]any, opts RenderOptions) (*Result, error) { if vars == nil { vars = map[string]any{} } - st := &state{ctx: ctx, vars: vars, maxSteps: opts.MaxSteps} + st := newState(ctx, vars, opts.MaxSteps, &budget{}) if st.maxSteps <= 0 { st.maxSteps = DefaultMaxSteps } @@ -75,9 +76,14 @@ func (*returnSignal) Error() string { return "vtl: #return" } type state struct { ctx context.Context vars map[string]any - out strings.Builder + out boundedWriter steps int maxSteps int + mem *budget +} + +func newState(ctx context.Context, vars map[string]any, maxSteps int, mem *budget) *state { + return &state{ctx: ctx, vars: vars, maxSteps: maxSteps, mem: mem, out: boundedWriter{limit: MaxOutputBytes}} } func (s *state) step() error { @@ -112,14 +118,9 @@ func (s *state) run(body []node) error { func (s *state) exec(n node) error { switch t := n.(type) { case *textNode: - s.out.WriteString(t.text) + return s.execText(t) case *refNode: - v, err := s.evalRef(t.ref) - if err != nil { - return err - } - - s.out.WriteString(Stringify(v)) + return s.execRef(t) case *setNode: return s.execSet(t) case *ifNode: @@ -137,6 +138,30 @@ func (s *state) exec(n node) error { return nil } +func (s *state) execText(n *textNode) error { + for _, part := range n.parts { + if err := s.out.WriteString(part); err != nil { + return err + } + } + + return nil +} + +func (s *state) execRef(n *refNode) error { + v, err := s.evalRef(n.ref) + if err != nil { + return err + } + + text, err := format(v, s.out.limit-s.out.Len()) + if err != nil { + return err + } + + return s.out.WriteString(text) +} + func (s *state) execSet(n *setNode) error { v, err := s.eval(n.value) if err != nil { @@ -159,7 +184,7 @@ func (s *state) execSet(n *setNode) error { switch last.kind { case accProperty: if m, ok := parent.(*Map); ok { - m.Put(last.name, v) + return s.put(m, last.name, v) } case accIndex: idx, err := s.eval(last.index) @@ -167,7 +192,7 @@ func (s *state) execSet(n *setNode) error { return err } - setIndex(parent, idx, v) + return s.setIndex(parent, idx, v) case accMethod: return errorf("cannot #set a method call") } @@ -175,15 +200,49 @@ func (s *state) execSet(n *setNode) error { return nil } -func setIndex(target, idx, v any) { +// put stores v in m, charging a new entry against the memory budget. +func (s *state) put(m *Map, key string, v any) error { + if _, exists := m.Get(key); !exists { + if err := s.mem.charge(slotBytes + len(key)); err != nil { + return err + } + } + + m.Put(key, v) + + return nil +} + +func (s *state) setIndex(target, idx, v any) error { switch t := target.(type) { case *Map: - t.Put(Stringify(idx), v) + key, err := s.key(idx) + if err != nil { + return err + } + + return s.put(t, key, v) case *List: if i, ok := toInt(idx); ok && i >= 0 && i < len(t.Items) { t.Items[i] = v } } + + return nil +} + +// key renders a value used as a map key. +func (s *state) key(v any) (string, error) { + if k, ok := v.(string); ok { + return k, nil + } + + k, err := format(v, MaxOutputBytes) + if err != nil { + return "", err + } + + return k, s.mem.charge(len(k)) } func (s *state) execIf(n *ifNode) error { @@ -207,10 +266,7 @@ func (s *state) execForeach(n *foreachNode) error { return err } - items := iterItems(src) - if len(items) > MaxForeachIterations { - items = items[:MaxForeachIterations] - } + items := iterItems(src, MaxForeachIterations) prevVar, hadVar := s.vars[n.varName] prevLoop, hadLoop := s.vars["foreach"] @@ -224,10 +280,7 @@ func (s *state) execForeach(n *foreachNode) error { for i, it := range items { s.vars[n.varName] = it - s.vars["foreach"] = MapOf( - "index", int64(i), "count", int64(i+1), "hasNext", i < len(items)-1, - "first", i == 0, "last", i == len(items)-1, - ) + s.vars["foreach"] = loopInfo(i, len(items)) s.vars["velocityCount"] = int64(i + 1) err := s.run(n.body) @@ -243,6 +296,18 @@ func (s *state) execForeach(n *foreachNode) error { return nil } +// loopInfo is the $foreach object for iteration i of n. +func loopInfo(i, n int) *Map { + m := NewMap() + m.Put("index", int64(i)) + m.Put("count", int64(i+1)) + m.Put("hasNext", i < n-1) + m.Put("first", i == 0) + m.Put("last", i == n-1) + + return m +} + func restoreVar(vars map[string]any, name string, prev any, had bool) { if had { vars[name] = prev @@ -251,15 +316,17 @@ func restoreVar(vars map[string]any, name string, prev any, had bool) { } } -// iterItems returns what #foreach walks: a list's items, a map's values, or -// nothing. -func iterItems(v any) []any { +// iterItems returns up to limit of the items #foreach walks: a list's items, +// a map's values, or nothing. +func iterItems(v any, limit int) []any { switch t := v.(type) { case *List: - return append([]any(nil), t.Items...) + return append([]any(nil), t.Items[:min(limit, len(t.Items))]...) case *Map: - out := make([]any, 0, t.Len()) - for _, k := range t.keys { + keys := t.keys[:min(limit, len(t.keys))] + out := make([]any, 0, len(keys)) + + for _, k := range keys { out = append(out, t.vals[k]) } @@ -317,9 +384,11 @@ func (s *state) evalCollection(e expr) (any, error) { } // evalInterpolated renders a double-quoted string as a template sharing the -// caller's variables and step budget. +// caller's variables, step budget and memory budget. func (s *state) evalInterpolated(t *interpolated) (any, error) { - sub := &state{ctx: s.ctx, vars: s.vars, steps: s.steps, maxSteps: s.maxSteps} + sub := newState(s.ctx, s.vars, s.maxSteps, s.mem) + sub.steps = s.steps + err := sub.run(t.body) s.steps = sub.steps @@ -327,10 +396,16 @@ func (s *state) evalInterpolated(t *interpolated) (any, error) { return nil, err } - return sub.out.String(), nil + out := sub.out.String() + + return out, s.mem.charge(len(out)) } func (s *state) evalList(t *listExpr) (any, error) { + if err := s.mem.charge(len(t.items) * slotBytes); err != nil { + return nil, err + } + l := NewList() for _, it := range t.items { @@ -378,7 +453,7 @@ func (s *state) evalRange(t *rangeExpr) (any, error) { } } - return l, nil + return l, s.mem.charge(len(l.Items) * slotBytes) } func (s *state) evalMap(t *mapExpr) (any, error) { @@ -395,7 +470,14 @@ func (s *state) evalMap(t *mapExpr) (any, error) { return nil, err } - m.Put(Stringify(k), v) + key, err := s.key(k) + if err != nil { + return nil, err + } + + if err := s.put(m, key, v); err != nil { + return nil, err + } } return m, nil @@ -451,7 +533,7 @@ func (s *state) evalBinary(t *binaryExpr) (any, error) { case opLt, opGt, opLe, opGe: return compare(t.op, l, r), nil default: - return arith(t.op, l, r), nil + return s.arith(t.op, l, r) } } @@ -504,7 +586,14 @@ func (s *state) access(cur any, a accessor) (any, error) { args[i] = v } - return callMethod(cur, a.name, args) + r, err := callMethod(s.mem, cur, a.name, args) + if err != nil { + return nil, err + } + + // A single method call can be slow on a large string, so the deadline + // is checked after every one. + return r, s.ctx.Err() } return nil, nil @@ -524,7 +613,7 @@ func property(v any, name string) any { // Bean-style getters: $list.empty, $str.empty. if name == "empty" { - if r, err := callMethod(v, mIsEmpty, nil); err == nil { + if r, err := callMethod(&budget{}, v, mIsEmpty, nil); err == nil { return r } } @@ -605,8 +694,36 @@ func equal(l, r any) bool { return ok && a == b } - // Different types compare by their string form, as Velocity does. - return Stringify(l) == Stringify(r) + return equalForms(l, r) +} + +// equalForms compares values of different types, and collections, by their +// string form. Forms too large or deep to render are unequal. +func equalForms(l, r any) bool { + if sameCollection(l, r) { + return true + } + + ls, lerr := format(l, MaxOutputBytes) + rs, rerr := format(r, MaxOutputBytes) + + return lerr == nil && rerr == nil && ls == rs +} + +// sameCollection reports whether l and r are the same list or map. +func sameCollection(l, r any) bool { + switch lt := l.(type) { + case *List: + rt, ok := r.(*List) + + return ok && lt == rt + case *Map: + rt, ok := r.(*Map) + + return ok && lt == rt + default: + return false + } } func compare(op string, l, r any) bool { @@ -654,13 +771,13 @@ func cmpFloat(a, b float64) int { // arith applies + - * / %. A string operand of + concatenates; integer // operands keep integer arithmetic; a division by zero is null. -func arith(op string, l, r any) any { +func (s *state) arith(op string, l, r any) (any, error) { if op == opAdd { _, ls := l.(string) _, rs := r.(string) if ls || rs { - return Stringify(l) + Stringify(r) + return s.concat(l, r) } } @@ -668,17 +785,36 @@ func arith(op string, l, r any) any { ri, rInt := r.(int64) if lInt && rInt { - return intArith(op, li, ri) + return intArith(op, li, ri), nil } a, okA := toFloat(l) b, okB := toFloat(r) if !okA || !okB { - return nil + return nil, nil + } + + return floatArith(op, a, b), nil +} + +// concat joins two values as strings within the output and memory limits. +func (s *state) concat(l, r any) (any, error) { + ls, err := format(l, MaxOutputBytes) + if err != nil { + return nil, err + } + + rs, err := format(r, MaxOutputBytes-len(ls)) + if err != nil { + return nil, err + } + + if err := s.mem.charge(len(ls) + len(rs)); err != nil { + return nil, err } - return floatArith(op, a, b) + return ls + rs, nil } func floatArith(op string, a, b float64) any { diff --git a/internal/vtl/fuzz_test.go b/internal/vtl/fuzz_test.go new file mode 100644 index 000000000..963aa46fb --- /dev/null +++ b/internal/vtl/fuzz_test.go @@ -0,0 +1,79 @@ +package vtl + +import ( + "context" + "runtime" + "testing" + "time" +) + +// heapGuard is the most live heap one fuzz iteration may leave behind. +const heapGuard = 1 << 30 + +// FuzzParseRender feeds arbitrary templates through the parser and the +// evaluator. Neither may panic, overflow the stack, run past its deadline by +// much, produce more than MaxOutputBytes, or hold on to a runaway heap. +func FuzzParseRender(f *testing.F) { + for _, seed := range []string{ + `#set($l = [])#set($x = $l.add($l))$l`, + `#set($m = {})#set($x = $m.put("a", $m))$m $m.equals($m)`, + `#set($s = "ab")#foreach($i in [1..1000])#set($s = "$s$s")#end$s`, + `#set($l = [1])#foreach($i in [1..1000])#set($x = $l.addAll($l))#end`, + `#foreach($i in [1..9])#if($i % 2 == 0)$i#elseif($i > 5)x#{else}-#end#end`, + `$input.body $util.parseJson('{"a":[1,2]}').a[1] ${x.y} $!z #* c *# ## c`, + `#set($s = "a")$s.replaceAll("(a)", "$1$1").split("a") $s.matches("[a-")`, + `#if((((1 + 2) * 3) > 4 && !false || $nope))ok#end#stop after`, + `#set($a = [])#foreach($i in [1..1000])#set($a = [$a])#end$a`, + } { + f.Add(seed) + } + + f.Fuzz(func(t *testing.T, src string) { + tmpl, err := Parse(src) + if err != nil { + return + } + + ctx, cancel := context.WithTimeout(context.Background(), 500*time.Millisecond) + defer cancel() + + start := time.Now() + res, err := tmpl.Render(ctx, map[string]any{"x": NewMap(), "util": fuzzHost{}}, RenderOptions{}) + + if d := time.Since(start); d > 5*time.Second { + t.Fatalf("render took %v", d) + } + + if err == nil && len(res.Output) > MaxOutputBytes { + t.Fatalf("output %d bytes", len(res.Output)) + } + + var ms runtime.MemStats + + runtime.ReadMemStats(&ms) + + if ms.HeapAlloc > heapGuard { + runtime.GC() + runtime.ReadMemStats(&ms) + + if ms.HeapAlloc > heapGuard { + t.Fatalf("live heap %d bytes after render", ms.HeapAlloc) + } + } + }) +} + +// fuzzHost is a host object with a JSON parser, like $util. +type fuzzHost struct{} + +func (fuzzHost) Get(string) (any, bool) { return nil, false } + +func (fuzzHost) Call(name string, args []any) (any, bool, error) { + if name != "parseJson" || len(args) != 1 { + return nil, false, nil + } + + v, err := ParseJSON(Stringify(args[0])) + + return v, true, err +} diff --git a/internal/vtl/limits.go b/internal/vtl/limits.go new file mode 100644 index 000000000..db1cb1713 --- /dev/null +++ b/internal/vtl/limits.go @@ -0,0 +1,112 @@ +package vtl + +import ( + "errors" + "strings" +) + +// Size and depth limits. A template that hits one fails with an error instead +// of growing without bound. +const ( + // MaxTemplateBytes is the largest template Parse accepts: API Gateway's + // 300 KB mapping-template quota. + MaxTemplateBytes = 300 << 10 + // MaxOutputBytes caps a render's output and any single string value: API + // Gateway's 10 MB payload quota. + MaxOutputBytes = 10 << 20 + // MaxAllocBytes caps the strings and collection entries one render may + // create, so a loop that keeps doubling a value fails fast. + MaxAllocBytes = 64 << 20 + // MaxTemplateDepth caps how deeply directives, expressions and string + // interpolations may nest. + MaxTemplateDepth = 100 + // MaxValueDepth caps how deeply printing, encoding, comparing and JSON + // parsing descend into nested lists and maps. + MaxValueDepth = 1000 + // maxOperatorChain caps a run of binary operators, which the evaluator + // walks recursively. + maxOperatorChain = 1000 + // slotBytes is what one list or map entry is charged against + // MaxAllocBytes: the 16-byte interface plus room for slice growth. + slotBytes = 32 +) + +// Limit errors. +var ( + ErrOutputLimit = errors.New("vtl: output exceeds the maximum size") + ErrMemoryLimit = errors.New("vtl: template exceeded its memory budget") + ErrDepthLimit = errors.New("vtl: value nesting exceeds the maximum depth") + ErrCyclicValue = errors.New("vtl: cannot encode a value that contains itself") +) + +// budget tracks the bytes a render has created. +type budget struct { + used int +} + +func (b *budget) charge(n int) error { + b.used += n + if b.used > MaxAllocBytes { + return ErrMemoryLimit + } + + return nil +} + +// sizeOf estimates the bytes a value holds, stopping once it passes limit. +// Shared and cyclic parts may be counted more than once, which only makes the +// estimate conservative. +func sizeOf(v any, limit int) int { + total := 0 + stack := []any{v} + visits := 0 + + for len(stack) > 0 && total <= limit && visits <= MaxAllocBytes/slotBytes { + cur := stack[len(stack)-1] + stack = stack[:len(stack)-1] + visits++ + + switch t := cur.(type) { + case string: + total += len(t) + case *List: + total += len(t.Items) * slotBytes + stack = append(stack, t.Items...) + case *Map: + total += t.Len() * slotBytes + + for _, k := range t.keys { + total += len(k) + stack = append(stack, t.vals[k]) + } + default: + total += slotBytes + } + } + + return total +} + +// boundedWriter is a strings.Builder that refuses to grow past limit. +type boundedWriter struct { + b strings.Builder + limit int +} + +func (w *boundedWriter) WriteString(s string) error { + if w.b.Len()+len(s) > w.limit { + return ErrOutputLimit + } + + w.b.WriteString(s) + + return nil +} + +func (w *boundedWriter) WriteByte(c byte) error { + return w.WriteString(string(c)) +} + +func (w *boundedWriter) String() string { return w.b.String() } + +func (w *boundedWriter) Len() int { return w.b.Len() } diff --git a/internal/vtl/limits_test.go b/internal/vtl/limits_test.go new file mode 100644 index 000000000..80535d8a9 --- /dev/null +++ b/internal/vtl/limits_test.go @@ -0,0 +1,164 @@ +package vtl + +import ( + "context" + "errors" + "strings" + "testing" + "time" +) + +func mustJSON(t *testing.T, v any) string { + t.Helper() + + s, err := ToJSON(v) + if err != nil { + t.Fatalf("ToJSON: %v", err) + } + + return s +} + +func renderErr(t *testing.T, src string, vars map[string]any) error { + t.Helper() + + tmpl, err := Parse(src) + if err != nil { + return err + } + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + _, err = tmpl.Render(ctx, vars, RenderOptions{}) + + return err +} + +func TestSelfReferencingValues(t *testing.T) { + cases := []struct { + name, src, want string + }{ + {"list prints itself", `#set($l = [1])#set($x = $l.add($l))$l`, "[1, (this Collection)]"}, + {"map prints itself", `#set($m = {"a": 1})#set($x = $m.put("self", $m))$m`, "{a=1, self=(this Map)}"}, + {"indirect cycle", `#set($a = [])#set($b = [$a])#set($x = $a.add($b))$a`, "[[(this Collection)]]"}, + {"cyclic equality", `#set($l = [])#set($x = $l.add($l))#if($l == $l)same#end#if($l.contains($l))has#end`, "samehas"}, + {"cyclic concat", `#set($l = [])#set($x = $l.add($l))#set($s = "v=$l")$s`, "v=[(this Collection)]"}, + } + + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + if got := render(t, c.src, nil); got != c.want { + t.Fatalf("got %q, want %q", got, c.want) + } + }) + } + + l := NewList() + l.Items = append(l.Items, l) + + if _, err := ToJSON(l); !errors.Is(err, ErrCyclicValue) { + t.Fatalf("ToJSON(cyclic) err = %v", err) + } + + if s := Stringify(l); s != "[(this Collection)]" { + t.Fatalf("Stringify(cyclic) = %q", s) + } +} + +func TestDeepValuesAreBounded(t *testing.T) { + // Each iteration wraps the list one level deeper. + src := `#set($a = [])#foreach($i in [1..1000])#set($a = [$a])#end#foreach($i in [1..1000])#set($a = [$a])#end$a` + if err := renderErr(t, src, nil); !errors.Is(err, ErrDepthLimit) { + t.Fatalf("deep print err = %v", err) + } + + deep := strings.Repeat("[", MaxValueDepth+1) + strings.Repeat("]", MaxValueDepth+1) + if _, err := ParseJSON(deep); !errors.Is(err, ErrDepthLimit) { + t.Fatalf("deep JSON err = %v", err) + } + + ok := strings.Repeat("[", 50) + strings.Repeat("]", 50) + if _, err := ParseJSON(ok); err != nil { + t.Fatalf("50-deep JSON: %v", err) + } +} + +func TestGrowthIsBounded(t *testing.T) { + cases := []struct { + name, src string + want error + }{ + {"string doubling", `#set($s = "ab")#foreach($i in [1..100])#set($s = "$s$s")#end`, ErrMemoryLimit}, + {"concat doubling", `#set($s = "ab")#foreach($i in [1..100])#set($s = $s + $s)#end`, ErrOutputLimit}, + {"list doubling", `#set($l = [1])#foreach($i in [1..100])#set($x = $l.addAll($l))#end`, ErrMemoryLimit}, + {"map growth", `#set($m = {})#foreach($i in [1..1000])#foreach($j in [1..1000])#set($m["$i-$j"] = $i)#end#end`, ErrStepBudget}, + {"replace blowup", `#set($s = "aaaaaaaaaa")#foreach($i in [1..40])#set($s = $s.replace("a", "aa"))#end`, ErrOutputLimit}, + {"literal replaceAll blowup", `#set($s = "aaaaaaaaaa")#foreach($i in [1..40])#set($s = $s.replaceAll("a", "aa"))#end`, ErrOutputLimit}, + {"regex blowup", `#set($s = "aaaaaaaaaa")#foreach($i in [1..40])#set($s = $s.replaceAll("[a]", "$0$0"))#end`, ErrOutputLimit}, + {"output flood", `#set($s = "0123456789")#foreach($i in [1..20])#set($s = "$s$s")#end#foreach($i in [1..1000])$s#end`, ErrOutputLimit}, + } + + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + start := time.Now() + + err := renderErr(t, c.src, nil) + if !errors.Is(err, c.want) && !isLimit(err) { + t.Fatalf("err = %v, want a limit error", err) + } + + // A render stops at its deadline (5s here) plus at most one slow + // call; the bound is loose so a loaded -race run stays green. + if d := time.Since(start); d > 20*time.Second { + t.Fatalf("took %v", d) + } + }) + } +} + +func TestParserLimits(t *testing.T) { + if _, err := Parse(strings.Repeat("a", MaxTemplateBytes+1)); err == nil { + t.Fatal("oversized template accepted") + } + + for name, src := range map[string]string{ + "nested parens": "#set($x = " + strings.Repeat("(", 200) + "1" + strings.Repeat(")", 200) + ")", + "nested ifs": strings.Repeat("#if(true)", 200) + strings.Repeat("#end", 200), + "nested not": "#set($x = " + strings.Repeat("!", 200) + "true)", + "operator chain": "#set($x = 1" + strings.Repeat(" + 1", maxOperatorChain+1) + ")", + } { + if _, err := Parse(src); err == nil { + t.Errorf("%s accepted", name) + } + } + + if got := render(t, "#set($x = 1"+strings.Repeat(" + 1", 500)+")$x", nil); got != "501" { + t.Fatalf("500-operator chain = %s", got) + } +} + +// TestParseIsLinear guards against quadratic parsing: a template of many short +// directives and literal dollars parses in well under a second. +func TestParseIsLinear(t *testing.T) { + var b strings.Builder + + for b.Len() < MaxTemplateBytes-64 { + b.WriteString(" #set($a = 1)\n$ $ #* c *#\n") + } + + start := time.Now() + + if _, err := Parse(b.String()); err != nil { + t.Fatal(err) + } + + if d := time.Since(start); d > time.Second { + t.Fatalf("parse took %v", d) + } +} + +func isLimit(err error) bool { + return errors.Is(err, ErrMemoryLimit) || errors.Is(err, ErrOutputLimit) || errors.Is(err, ErrStepBudget) || + errors.Is(err, context.DeadlineExceeded) +} diff --git a/internal/vtl/methods.go b/internal/vtl/methods.go index 6475baabb..552f46af2 100644 --- a/internal/vtl/methods.go +++ b/internal/vtl/methods.go @@ -15,6 +15,9 @@ const ( mRemove = "remove" mSet = "set" mPut = "put" + mKeySet = "keySet" + mValues = "values" + mEntrySet = "entrySet" ) // pairArgs is the argument count of a two-argument method such as put or set. @@ -69,61 +72,137 @@ var ( "putAll": mapPutAll, "containsKey": mapContainsKey, mRemove: func(m *Map, args []any) any { return m.Remove(Stringify(firstArg(args))) }, - "keySet": func(m *Map, _ []any) any { return stringList(m.keys) }, - "values": func(m *Map, _ []any) any { return NewList(iterItems(m)...) }, - "entrySet": mapEntrySet, + mKeySet: func(m *Map, _ []any) any { return stringList(m.keys) }, + mValues: func(m *Map, _ []any) any { return NewList(iterItems(m, m.Len())...) }, + mEntrySet: mapEntrySet, mSize: func(m *Map, _ []any) any { return int64(m.Len()) }, mIsEmpty: func(m *Map, _ []any) any { return m.Len() == 0 }, } ) +// creatingMethods are the list and map methods that return a new collection. +var creatingMethods = map[string]bool{mKeySet: true, mValues: true, mEntrySet: true} //nolint:gochecknoglobals // read-only + // callMethod dispatches a method call to the bridge for strings, lists and -// maps, or to a host Object. An unknown method yields null, as Velocity does -// when no method matches. -func callMethod(v any, name string, args []any) (any, error) { +// maps, or to a host Object, charging what the call creates against mem. An +// unknown method yields null, as Velocity does when no method matches. +func callMethod(mem *budget, v any, name string, args []any) (any, error) { switch { case name == "toString" && len(args) == 0: - return Stringify(v), nil + return chargeString(mem, v) case name == "equals" && len(args) == 1: return equal(v, args[0]), nil } switch t := v.(type) { case string: - if fn, ok := stringMethods[name]; ok { - return fn(t, args) - } + return callString(mem, t, name, args) case Object: - return callObject(t, name, args) + return callObject(mem, t, name, args) default: - return callCollection(v, name, args), nil + return callCollection(mem, v, name, args) + } +} + +func chargeString(mem *budget, v any) (any, error) { + s, err := format(v, MaxOutputBytes) + if err != nil { + return nil, err } - return nil, nil + return s, mem.charge(len(s)) } -func callCollection(v any, name string, args []any) any { +func callString(mem *budget, s, name string, args []any) (any, error) { + fn, ok := stringMethods[name] + if !ok { + return nil, nil + } + + r, err := fn(s, args) + if err != nil { + return nil, err + } + + switch t := r.(type) { + case string: + if len(t) > MaxOutputBytes { + return nil, ErrOutputLimit + } + + return t, mem.charge(len(t)) + case *List: + return t, mem.charge(len(t.Items) * slotBytes) + default: + return r, nil + } +} + +// callCollection runs a list or map method, charging any growth of the +// receiver and any new collection it returns. +func callCollection(mem *budget, v any, name string, args []any) (any, error) { + var ( + r any + before = collectionLen(v) + ) + switch t := v.(type) { case *List: - if fn, ok := listMethods[name]; ok { - return fn(t, args) + fn, ok := listMethods[name] + if !ok { + return nil, nil } + + r = fn(t, args) case *Map: - if fn, ok := mapMethods[name]; ok { - return fn(t, args) + fn, ok := mapMethods[name] + if !ok { + return nil, nil } + + r = fn(t, args) + default: + return nil, nil } - return nil + grown := collectionLen(v) - before + if creatingMethods[name] { + grown += collectionLen(r) + } + + if grown > 0 { + if err := mem.charge(grown * slotBytes); err != nil { + return nil, err + } + } + + return r, nil } -func callObject(o Object, name string, args []any) (any, error) { +func collectionLen(v any) int { + switch t := v.(type) { + case *List: + return len(t.Items) + case *Map: + return t.Len() + default: + return 0 + } +} + +// callObject calls a host method. Its result is charged by size, since the +// host may build a new value from template-controlled input. +func callObject(mem *budget, o Object, name string, args []any) (any, error) { r, ok, err := o.Call(name, args) if err != nil || !ok { return nil, err } - return r, nil + if s, isStr := r.(string); isStr && len(s) > MaxOutputBytes { + return nil, ErrOutputLimit + } + + return r, mem.charge(sizeOf(r, MaxAllocBytes)) } // strArg returns argument i as a string; a non-string is stringified. @@ -181,6 +260,25 @@ func strReplace(s string, args []any) (any, error) { return nil, nil } + return literalReplace(s, from, to, true) +} + +// literalReplace replaces from with to (every occurrence, or the first), +// checking the result size before building it. +func literalReplace(s, from, to string, all bool) (any, error) { + n := 1 + if all { + n = strings.Count(s, from) + } + + if strings.Contains(s, from) && len(s)+n*(len(to)-len(from)) > MaxOutputBytes { + return nil, ErrOutputLimit + } + + if !all { + return strings.Replace(s, from, to, 1), nil + } + return strings.ReplaceAll(s, from, to), nil } @@ -240,21 +338,43 @@ func regexReplace(s string, all bool, args []any) (any, error) { pattern, _ := strArg(args, 0) repl, _ := strArg(args, 1) + // A literal pattern and replacement need no regex engine. + if regexp.QuoteMeta(pattern) == pattern && !strings.ContainsAny(repl, `$\`) && pattern != "" { + return literalReplace(s, pattern, repl, all) + } + re, err := compile(pattern) if err != nil { return nil, err } + limit := 1 if all { - return re.ReplaceAllString(s, repl), nil + limit = -1 + } + + // Build the result match by match so it can stop at the size limit. + var ( + out []byte + last int + ) + + for _, loc := range re.FindAllStringSubmatchIndex(s, limit) { + out = append(out, s[last:loc[0]]...) + out = re.ExpandString(out, repl, s, loc) + last = loc[1] + + if len(out) > MaxOutputBytes { + return nil, ErrOutputLimit + } } - loc := re.FindStringSubmatchIndex(s) - if loc == nil { - return s, nil + out = append(out, s[last:]...) + if len(out) > MaxOutputBytes { + return nil, ErrOutputLimit } - return s[:loc[0]] + string(re.ExpandString(nil, repl, s, loc)) + s[loc[1]:], nil + return string(out), nil } // strSplit follows Java's String.split: the argument is a regex and trailing diff --git a/internal/vtl/parse.go b/internal/vtl/parse.go index c147cf6c7..eb5d9518c 100644 --- a/internal/vtl/parse.go +++ b/internal/vtl/parse.go @@ -21,8 +21,13 @@ func (e *ParseError) Error() string { return fmt.Sprintf("vtl: parse error at offset %d: %s", e.Pos, e.Msg) } -// Parse parses src into a Template. +// Parse parses src into a Template. A template larger than MaxTemplateBytes +// or nested deeper than MaxTemplateDepth is rejected. func Parse(src string) (*Template, error) { + if len(src) > MaxTemplateBytes { + return nil, &ParseError{Msg: fmt.Sprintf("template is %d bytes, more than the %d byte limit", len(src), MaxTemplateBytes)} + } + p := &parser{src: src} body, term, err := p.parseBlock() @@ -40,8 +45,25 @@ func Parse(src string) (*Template, error) { type parser struct { src string pos int + // depth is the current nesting of blocks, expressions and interpolated + // strings. + depth int + // dirStart is where the most recent directive began. + dirStart int } +// enter descends one nesting level, failing past MaxTemplateDepth. +func (p *parser) enter() error { + p.depth++ + if p.depth > MaxTemplateDepth { + return p.errf("template nests deeper than %d levels", MaxTemplateDepth) + } + + return nil +} + +func (p *parser) leave() { p.depth-- } + func (p *parser) errf(format string, args ...any) error { return &ParseError{Pos: p.pos, Msg: fmt.Sprintf(format, args...)} } @@ -73,6 +95,12 @@ var unsupportedDirectives = map[string]bool{ //nolint:gochecknoglobals // read-o // parseBlock parses nodes until EOF or a block terminator (#else, #elseif, // #end), which it returns. For #elseif the parser is left at its condition. func (p *parser) parseBlock() ([]node, string, error) { + if err := p.enter(); err != nil { + return nil, "", err + } + + defer p.leave() + var nodes []node for p.pos < len(p.src) { @@ -120,13 +148,13 @@ func appendText(nodes []node, s string) []node { if n := len(nodes); n > 0 { if t, ok := nodes[n-1].(*textNode); ok { - t.text += s + t.parts = append(t.parts, s) return nodes } } - return append(nodes, &textNode{text: s}) + return append(nodes, &textNode{parts: []string{s}}) } // parseEscape handles a backslash: \$ and \# print the next character @@ -199,6 +227,8 @@ func (p *parser) parseHash(nodes []node) ([]node, string, error) { return appendText(nodes, rest[3:end]), "", nil } + p.dirStart = start + name, braced := p.directiveName() if unsupportedDirectives[name] { return nil, "", p.errf("#%s is not supported", name) @@ -304,31 +334,63 @@ func (p *parser) gobble(nodes []node, start int) []node { if n := len(nodes); n > 0 { if t, ok := nodes[n-1].(*textNode); ok { - t.text = strings.TrimRight(t.text, " \t") + trimIndent(t) + + if len(t.parts) == 0 { + nodes = nodes[:n-1] + } } } return nodes } +// trimIndent drops trailing spaces and tabs from a text node. +func trimIndent(t *textNode) { + for len(t.parts) > 0 { + last := len(t.parts) - 1 + trimmed := strings.TrimRight(t.parts[last], " \t") + + if trimmed != "" { + t.parts[last] = trimmed + + return + } + + t.parts = t.parts[:last] + } +} + // blankLineEnd reports whether the directive spanning start..p.pos is alone on -// its line, and returns the offset just past that line's newline. +// its line, and returns the offset just past that line's newline. It only +// scans the whitespace around the directive, so parsing stays linear. func (p *parser) blankLineEnd(start int) (int, bool) { - lineStart := strings.LastIndexByte(p.src[:start], '\n') + 1 - if strings.TrimLeft(p.src[lineStart:start], " \t") != "" { - return 0, false + i := start + for i > 0 && isBlank(p.src[i-1]) { + i-- } - rest := p.src[p.pos:] - nl := strings.IndexByte(rest, '\n') + if i > 0 && p.src[i-1] != '\n' { + return 0, false + } - if nl < 0 { - return len(p.src), strings.TrimLeft(rest, " \t\r") == "" + j := p.pos + for j < len(p.src) && (isBlank(p.src[j]) || p.src[j] == '\r') { + j++ } - return p.pos + nl + 1, strings.TrimLeft(rest[:nl], " \t\r") == "" + switch { + case j == len(p.src): + return j, true + case p.src[j] == '\n': + return j + 1, true + default: + return 0, false + } } +func isBlank(c byte) bool { return c == ' ' || c == '\t' } + func (p *parser) parseSet() (node, error) { if err := p.expectOpenParen(); err != nil { return nil, err @@ -399,7 +461,7 @@ func (p *parser) parseIf(nodes []node, start int) ([]node, string, error) { switch term { case dirElseIf: - elseStart := strings.LastIndex(p.src[:p.pos], "#") + elseStart := p.dirStart if cond, err = p.parseCondition(); err != nil { return nil, "", err @@ -644,7 +706,15 @@ func (p *parser) parseArgs(closer byte) ([]expr, error) { // Expression parsing, lowest precedence first. -func (p *parser) parseExpr() (expr, error) { return p.parseOr() } +func (p *parser) parseExpr() (expr, error) { + if err := p.enter(); err != nil { + return nil, err + } + + defer p.leave() + + return p.parseOr() +} func (p *parser) parseOr() (expr, error) { return p.parseBinary(p.parseAnd, map[string]string{opOr: opOr, "or": opOr}) @@ -679,12 +749,16 @@ func (p *parser) parseBinary(next func() (expr, error), ops map[string]string) ( return nil, err } - for { + for chain := 0; ; chain++ { op, ok := p.matchOperator(ops) if !ok { return l, nil } + if chain >= maxOperatorChain { + return nil, p.errf("more than %d operators in one expression", maxOperatorChain) + } + r, err := next() if err != nil { return nil, err @@ -724,37 +798,44 @@ func (p *parser) matchOperator(ops map[string]string) (string, bool) { } func (p *parser) parseUnary() (expr, error) { + if err := p.enter(); err != nil { + return nil, err + } + + defer p.leave() + p.skipSpace() + op := p.unaryOperator() + if op == "" { + return p.parsePrimary() + } + + x, err := p.parseUnary() + if err != nil { + return nil, err + } + + return &unaryExpr{op: op, x: x}, nil +} + +// unaryOperator consumes a prefix ! / not (as opNot) or a minus before a +// non-number (as opSub), returning "" when there is none. +func (p *parser) unaryOperator() string { switch { case p.peek() == '!' && !strings.HasPrefix(p.src[p.pos:], opNe): p.pos++ - x, err := p.parseUnary() - if err != nil { - return nil, err - } - - return &unaryExpr{op: opNot, x: x}, nil + return opNot case p.consumeWord("not"): - x, err := p.parseUnary() - if err != nil { - return nil, err - } - - return &unaryExpr{op: opNot, x: x}, nil + return opNot case p.peek() == '-' && p.pos+1 < len(p.src) && !isDigit(p.src[p.pos+1]): p.pos++ - x, err := p.parseUnary() - if err != nil { - return nil, err - } - - return &unaryExpr{op: opSub, x: x}, nil + return opSub + default: + return "" } - - return p.parsePrimary() } func (p *parser) parsePrimary() (expr, error) { @@ -846,7 +927,7 @@ func (p *parser) parseDoubleQuoted() (expr, error) { return &literal{value: s}, nil } - sub := &parser{src: s} + sub := &parser{src: s, depth: p.depth} body, term, err := sub.parseBlock() if err != nil { diff --git a/internal/vtl/value.go b/internal/vtl/value.go index 6de90ae9b..763edb9fe 100644 --- a/internal/vtl/value.go +++ b/internal/vtl/value.go @@ -43,6 +43,9 @@ func (l *List) Index(i int) (any, bool) { return l.Items[i], true } +// Len implements jsonpath.Array. +func (l *List) Len() int { return len(l.Items) } + // Map is an insertion-ordered, mutable map value (Java's LinkedHashMap). type Map struct { keys []string @@ -52,17 +55,6 @@ type Map struct { // NewMap returns an empty map. func NewMap() *Map { return &Map{vals: map[string]any{}} } -// MapOf builds a map from alternating key/value arguments. -func MapOf(kv ...any) *Map { - m := NewMap() - - for i := 0; i+1 < len(kv); i += 2 { - m.Put(fmt.Sprint(kv[i]), kv[i+1]) - } - - return m -} - // StringMap converts a Go string map to a Map with keys in sorted order. func StringMap(in map[string]string) *Map { m := NewMap() @@ -116,19 +108,20 @@ func (m *Map) Remove(key string) any { return prev } -// Keys returns the keys in insertion order. +// Keys returns the keys in insertion order. It implements jsonpath.Object. func (m *Map) Keys() []string { return append([]string(nil), m.keys...) } // Len returns the number of entries. func (m *Map) Len() int { return len(m.keys) } // ParseJSON decodes a JSON document into template values, keeping object key -// order and decoding integral numbers as int64. +// order and decoding integral numbers as int64. Nesting deeper than +// MaxValueDepth is rejected. func ParseJSON(s string) (any, error) { dec := json.NewDecoder(strings.NewReader(s)) dec.UseNumber() - v, err := decodeValue(dec) + v, err := decodeValue(dec, 0) if err != nil { return nil, err } @@ -140,7 +133,7 @@ func ParseJSON(s string) (any, error) { return v, nil } -func decodeValue(dec *json.Decoder) (any, error) { +func decodeValue(dec *json.Decoder, depth int) (any, error) { tok, err := dec.Token() if err != nil { return nil, err @@ -148,12 +141,16 @@ func decodeValue(dec *json.Decoder) (any, error) { switch t := tok.(type) { case json.Delim: + if depth >= MaxValueDepth { + return nil, ErrDepthLimit + } + if t == '{' { - return decodeObject(dec) + return decodeObject(dec, depth+1) } if t == '[' { - return decodeArray(dec) + return decodeArray(dec, depth+1) } return nil, errorf("unexpected %q in JSON", t) @@ -164,7 +161,7 @@ func decodeValue(dec *json.Decoder) (any, error) { } } -func decodeObject(dec *json.Decoder) (any, error) { +func decodeObject(dec *json.Decoder, depth int) (any, error) { m := NewMap() for dec.More() { @@ -175,7 +172,7 @@ func decodeObject(dec *json.Decoder) (any, error) { key, _ := tok.(string) - v, err := decodeValue(dec) + v, err := decodeValue(dec, depth) if err != nil { return nil, err } @@ -190,11 +187,11 @@ func decodeObject(dec *json.Decoder) (any, error) { return m, nil } -func decodeArray(dec *json.Decoder) (any, error) { +func decodeArray(dec *json.Decoder, depth int) (any, error) { l := NewList() for dec.More() { - v, err := decodeValue(dec) + v, err := decodeValue(dec, depth) if err != nil { return nil, err } @@ -219,65 +216,246 @@ func jsonNumber(n json.Number) any { return f } -// ToJSON encodes a template value as JSON. Host objects encode as null. -func ToJSON(v any) string { - var b bytes.Buffer +// ToJSON encodes a template value as JSON. Host objects encode as null. A +// value that contains itself, nests deeper than MaxValueDepth or encodes to +// more than MaxOutputBytes is an error. +func ToJSON(v any) (string, error) { + f := &formatter{w: boundedWriter{limit: MaxOutputBytes}} + if err := f.json(v, 0); err != nil { + return "", err + } + + return f.w.String(), nil +} + +// Stringify renders a value the way Velocity prints it: nil as empty, lists +// as [a, b] and maps as {k=v, k2=v2}, as Java's toString does. A list or map +// that contains itself prints as "(this Collection)" or "(this Map)". Output +// past MaxOutputBytes or nesting past MaxValueDepth is cut off; use format +// where that must be an error. +func Stringify(v any) string { + s, _ := format(v, MaxOutputBytes) + + return s +} + +// format renders v like Stringify, failing once the result would pass limit +// bytes or nest past MaxValueDepth. +func format(v any, limit int) (string, error) { + if s, ok := v.(string); ok { + if len(s) > limit { + return "", ErrOutputLimit + } + + return s, nil + } + + f := &formatter{w: boundedWriter{limit: limit}} + err := f.str(v, 0) + + return f.w.String(), err +} + +// formatter writes values into a bounded buffer, tracking the lists and maps +// it is inside so a self-reference is detected. +type formatter struct { + w boundedWriter + inside map[any]bool +} + +func (f *formatter) enter(v any, depth int) (cyclic bool, err error) { + if depth >= MaxValueDepth { + return false, ErrDepthLimit + } - writeJSON(&b, v) + if f.inside[v] { + return true, nil + } + + if f.inside == nil { + f.inside = map[any]bool{} + } + + f.inside[v] = true - return b.String() + return false, nil } -func writeJSON(b *bytes.Buffer, v any) { +func (f *formatter) leave(v any) { delete(f.inside, v) } + +func (f *formatter) str(v any, depth int) error { + switch t := v.(type) { + case *List: + return f.strList(t, depth) + case *Map: + return f.strMap(t, depth) + default: + return f.w.WriteString(scalarString(v)) + } +} + +func (f *formatter) strList(l *List, depth int) error { + cyclic, err := f.enter(l, depth) + if err != nil { + return err + } + + if cyclic { + return f.w.WriteString("(this Collection)") + } + + defer f.leave(l) + + if err := f.w.WriteByte('['); err != nil { + return err + } + + for i, it := range l.Items { + if i > 0 { + if err := f.w.WriteString(", "); err != nil { + return err + } + } + + if err := f.str(it, depth+1); err != nil { + return err + } + } + + return f.w.WriteByte(']') +} + +func (f *formatter) strMap(m *Map, depth int) error { + cyclic, err := f.enter(m, depth) + if err != nil { + return err + } + + if cyclic { + return f.w.WriteString("(this Map)") + } + + defer f.leave(m) + + if err := f.w.WriteByte('{'); err != nil { + return err + } + + for i, k := range m.keys { + sep := k + "=" + if i > 0 { + sep = ", " + sep + } + + if err := f.w.WriteString(sep); err != nil { + return err + } + + if err := f.str(m.vals[k], depth+1); err != nil { + return err + } + } + + return f.w.WriteByte('}') +} + +func (f *formatter) json(v any, depth int) error { switch t := v.(type) { case nil: - b.WriteString("null") + return f.w.WriteString("null") case string: - writeJSONString(b, t) + return f.w.WriteString(jsonString(t)) case bool, int64, float64: - b.WriteString(Stringify(t)) + return f.w.WriteString(scalarString(t)) case *List: - b.WriteByte('[') + return f.jsonList(t, depth) + case *Map: + return f.jsonMap(t, depth) + default: + return f.w.WriteString("null") + } +} + +func (f *formatter) jsonList(l *List, depth int) error { + if err := f.enterJSON(l, depth); err != nil { + return err + } + + defer f.leave(l) + + if err := f.w.WriteByte('['); err != nil { + return err + } - for i, it := range t.Items { - if i > 0 { - b.WriteByte(',') + for i, it := range l.Items { + if i > 0 { + if err := f.w.WriteByte(','); err != nil { + return err } + } - writeJSON(b, it) + if err := f.json(it, depth+1); err != nil { + return err } + } - b.WriteByte(']') - case *Map: - b.WriteByte('{') + return f.w.WriteByte(']') +} - for i, k := range t.keys { - if i > 0 { - b.WriteByte(',') - } +func (f *formatter) jsonMap(m *Map, depth int) error { + if err := f.enterJSON(m, depth); err != nil { + return err + } + + defer f.leave(m) - writeJSONString(b, k) - b.WriteByte(':') - writeJSON(b, t.vals[k]) + if err := f.w.WriteByte('{'); err != nil { + return err + } + + for i, k := range m.keys { + key := jsonString(k) + ":" + if i > 0 { + key = "," + key } - b.WriteByte('}') - default: - b.WriteString("null") + if err := f.w.WriteString(key); err != nil { + return err + } + + if err := f.json(m.vals[k], depth+1); err != nil { + return err + } } + + return f.w.WriteByte('}') +} + +func (f *formatter) enterJSON(v any, depth int) error { + cyclic, err := f.enter(v, depth) + if err != nil { + return err + } + + if cyclic { + return ErrCyclicValue + } + + return nil } -func writeJSONString(b *bytes.Buffer, s string) { - enc := json.NewEncoder(b) +func jsonString(s string) string { + var b bytes.Buffer + + enc := json.NewEncoder(&b) enc.SetEscapeHTML(false) _ = enc.Encode(s) - b.Truncate(b.Len() - 1) // drop the encoder's trailing newline + return strings.TrimSuffix(b.String(), "\n") } -// Stringify renders a value the way Velocity prints it: nil as empty, lists -// as [a, b] and maps as {k=v, k2=v2}, as Java's toString does. -func Stringify(v any) string { +// scalarString prints a non-collection value. +func scalarString(v any) string { switch t := v.(type) { case nil: return "" @@ -289,28 +467,6 @@ func Stringify(v any) string { return strconv.FormatInt(t, 10) case float64: return formatFloat(t) - default: - return stringifyComposite(v) - } -} - -// stringifyComposite renders lists, maps and Stringers. -func stringifyComposite(v any) string { - switch t := v.(type) { - case *List: - parts := make([]string, len(t.Items)) - for i, it := range t.Items { - parts[i] = Stringify(it) - } - - return "[" + strings.Join(parts, ", ") + "]" - case *Map: - parts := make([]string, len(t.keys)) - for i, k := range t.keys { - parts[i] = k + "=" + Stringify(t.vals[k]) - } - - return "{" + strings.Join(parts, ", ") + "}" case fmt.Stringer: return t.String() default: diff --git a/internal/vtl/vtl_test.go b/internal/vtl/vtl_test.go index 9a303d86a..bf8d274e2 100644 --- a/internal/vtl/vtl_test.go +++ b/internal/vtl/vtl_test.go @@ -91,7 +91,7 @@ func TestReturnDirective(t *testing.T) { t.Fatal(err) } - if !res.Returned || res.Output != "before" || ToJSON(res.ReturnValue) != `{"a":1}` { + if !res.Returned || res.Output != "before" || mustJSON(t, res.ReturnValue) != `{"a":1}` { t.Fatalf("got %+v", res) } } @@ -168,7 +168,7 @@ func TestJSONRoundTripKeepsOrder(t *testing.T) { t.Fatal(err) } - if got := ToJSON(v); got != src { + if got := mustJSON(t, v); got != src { t.Fatalf("ToJSON = %s, want %s", got, src) } @@ -176,11 +176,11 @@ func TestJSONRoundTripKeepsOrder(t *testing.T) { t.Fatal("trailing data accepted") } - if got := ToJSON("<&>"); got != `"<&>"` { + if got := mustJSON(t, "<&>"); got != `"<&>"` { t.Fatalf("html escaped: %s", got) } - if !strings.Contains(ToJSON(StringMap(map[string]string{"b": "2", "a": "1"})), `{"a":"1","b":"2"}`) { + if !strings.Contains(mustJSON(t, StringMap(map[string]string{"b": "2", "a": "1"})), `{"a":"1","b":"2"}`) { t.Fatal("StringMap order") } } diff --git a/providers/aws/apigateway/apigateway.go b/providers/aws/apigateway/apigateway.go index 02b6013c8..cdeaab8f1 100644 --- a/providers/aws/apigateway/apigateway.go +++ b/providers/aws/apigateway/apigateway.go @@ -16,6 +16,7 @@ import ( "github.com/stackshy/cloudemu/v2/config" cerrors "github.com/stackshy/cloudemu/v2/errors" "github.com/stackshy/cloudemu/v2/internal/memstore" + "github.com/stackshy/cloudemu/v2/internal/vtl" "github.com/stackshy/cloudemu/v2/services/apigateway/driver" mondriver "github.com/stackshy/cloudemu/v2/services/monitoring/driver" ) @@ -83,13 +84,21 @@ type Mock struct { regionMu sync.RWMutex certs map[string]*driver.ClientCertificate account driver.Account + + // templates caches parsed mapping templates by source, so a deployed + // template is parsed once rather than on every invoke. + templates *vtl.Cache } +// templateCacheSize bounds the parsed mapping templates kept in memory. +const templateCacheSize = 512 + // New creates a new API Gateway mock. func New(opts *config.Options) *Mock { return &Mock{ apis: memstore.New[*apiData](), opts: opts, certs: map[string]*driver.ClientCertificate{}, account: defaultAccount(), + templates: vtl.NewCache(templateCacheSize), } } diff --git a/providers/aws/apigateway/mapping.go b/providers/aws/apigateway/mapping.go index 65c255580..5820fa056 100644 --- a/providers/aws/apigateway/mapping.go +++ b/providers/aws/apigateway/mapping.go @@ -40,6 +40,8 @@ type mappingContext struct { // context is the $context map. It is shared by the request and response // templates, so $context.responseOverride set in either survives. context *vtl.Map + // templates caches parsed templates; nil parses on every render. + templates *vtl.Cache } func newMappingContext(req *driver.ProxyRequest, route *resolvedRoute, account, reqID string, now time.Time) *mappingContext { @@ -83,16 +85,22 @@ func (mc *mappingContext) buildContext() *vtl.Map { return ctx } -// render evaluates a mapping template with body as $input's payload. +// render evaluates a mapping template with body as $input's payload. Parsing +// (cached per template source) and rendering share one deadline. func (mc *mappingContext) render(ctx context.Context, src, body string) (string, error) { - tmpl, err := vtl.Parse(src) + ctx, cancel := context.WithTimeout(ctx, templateTimeout) + defer cancel() + + parse := vtl.Parse + if mc.templates != nil { + parse = mc.templates.Parse + } + + tmpl, err := parse(src) if err != nil { return "", err } - ctx, cancel := context.WithTimeout(ctx, templateTimeout) - defer cancel() - vars := map[string]any{ "input": &inputObject{body: body, params: mc.params()}, "context": mc.context, @@ -179,7 +187,9 @@ func (in *inputObject) Call(name string, args []any) (res any, found bool, callE return nil, true, err } - return vtl.ToJSON(v), true, nil + s, err := vtl.ToJSON(v) + + return s, true, err case "params": if len(args) == 0 { return in.params, true, nil @@ -194,7 +204,8 @@ func (in *inputObject) Call(name string, args []any) (res any, found bool, callE } // path evaluates a JSONPath against the JSON body. An empty body is treated as -// an empty object, as API Gateway does. +// an empty object, as API Gateway does. A path with a wildcard or recursive +// descent returns the list of matches. func (in *inputObject) path(p string) (any, error) { if !in.done { in.done = true @@ -208,9 +219,20 @@ func (in *inputObject) path(p string) (any, error) { } } - v, _, err := jsonpath.Eval(p, in.parsed) + matches, indefinite, err := jsonpath.EvalAll(p, in.parsed) + if err != nil { + return nil, err + } + + if indefinite { + return vtl.NewList(matches...), nil + } - return v, err + if len(matches) == 0 { + return nil, nil + } + + return matches[0], nil } // param looks a name up in the path, querystring and header maps, in that @@ -263,7 +285,7 @@ func (utilObject) Call(name string, args []any) (res any, found bool, callErr er return v, true, nil case "urlEncode": - return url.QueryEscape(s), true, nil + return javaURLEncode(s), true, nil case "urlDecode": v, err := url.QueryUnescape(s) if err != nil { @@ -293,6 +315,13 @@ func stringArg(args []any) string { return vtl.Stringify(args[0]) } +// javaURLEncode matches Java's URLEncoder.encode with UTF-8, which +// $util.urlEncode uses: unlike Go's QueryEscape it leaves '*' alone and +// encodes '~'. +func javaURLEncode(s string) string { + return strings.NewReplacer("%2A", "*", "~", "%7E").Replace(url.QueryEscape(s)) +} + // escapeJavaScript matches Apache Commons StringEscapeUtils.escapeJavaScript, // which $util.escapeJavaScript uses: quotes, backslash and '/' are escaped, // control characters use their short or \uXXXX form, and non-ASCII characters @@ -329,7 +358,7 @@ func writeEscapedRune(b *strings.Builder, r rune) { lastASCII = 0x7f ) - if r >= firstPrintable && r < lastASCII { + if r >= firstPrintable && r <= lastASCII { b.WriteRune(r) return diff --git a/providers/aws/apigateway/mapping_limits_test.go b/providers/aws/apigateway/mapping_limits_test.go new file mode 100644 index 000000000..eb618d5fb --- /dev/null +++ b/providers/aws/apigateway/mapping_limits_test.go @@ -0,0 +1,102 @@ +package apigateway_test + +import ( + "strings" + "testing" + + "github.com/stackshy/cloudemu/v2/errors" + "github.com/stackshy/cloudemu/v2/services/apigateway/driver" +) + +// TestMockTemplateAbuseIsAnError covers templates that would otherwise recurse +// forever or grow without bound: each must end as a 500 from the gateway. +func TestMockTemplateAbuseIsAnError(t *testing.T) { + for name, tmpl := range map[string]string{ + "self list": `#set($l = [])#set($x = $l.add($l))$l`, + "self list json": `#set($l = [])#set($x = $l.add($l))$util.parseJson("[]")$input.json('$')`, + "string doubling": `#set($s = "ab")#foreach($i in [1..200])#set($s = "$s$s")#end$s`, + "list doubling": `#set($l = [1])#foreach($i in [1..200])#set($x = $l.addAll($l))#end$l.size()`, + } { + t.Run(name, func(t *testing.T) { + m := newMock(t) + apiID, _ := mockMethod(t, m, map[string]string{"application/json": `{"statusCode": 200}`}, "", tmpl) + + resp := invokeMock(t, m, apiID, driver.ProxyRequest{}) + + switch name { + case "self list": + if resp.StatusCode != 200 || resp.Body != "[(this Collection)]" { + t.Fatalf("got %d %q", resp.StatusCode, resp.Body) + } + case "self list json": + if resp.StatusCode != 200 || resp.Body != "[]{}" { + t.Fatalf("got %d %q", resp.StatusCode, resp.Body) + } + default: + if resp.StatusCode != 500 || resp.Headers["x-amzn-ErrorType"] != "InternalServerErrorException" { + t.Fatalf("got %d %.80q", resp.StatusCode, resp.Body) + } + } + }) + } +} + +func TestMockUtilEncodingAndPaths(t *testing.T) { + m := newMock(t) + tmpl := `$util.urlEncode('a~b*c d')|$util.escapeJavaScript("x` + "\x7f" + `y")|` + + `$input.json('$.items[*].id')|$input.json('$..id')|$input.path('$.items[0].id')` + apiID, _ := mockMethod(t, m, map[string]string{"application/json": `{"statusCode": 200}`}, "", tmpl) + + // The MOCK backend body is empty, so paths select from the response + // template's own (empty) input: the list forms render as empty lists. + resp := invokeMock(t, m, apiID, driver.ProxyRequest{}) + + want := "a%7Eb*c+d|x\x7fy|[]|[]|" + if resp.StatusCode != 200 || resp.Body != want { + t.Fatalf("got %d %q, want %q", resp.StatusCode, resp.Body, want) + } +} + +func TestMockRequestTemplatePaths(t *testing.T) { + m := newMock(t) + reqTmpl := `#set($ids = $input.path('$.items[*].id'))` + + `{"statusCode": #if($ids.size() == 2 && $input.json('$..id') == "[1,2]")201#{else}500#end}` + apiID, resID := mockMethod(t, m, map[string]string{"application/json": reqTmpl}, "", "") + + if _, err := m.PutMethodResponse(ctx(), apiID, resID, "GET", "201", driver.PutMethodResponseInput{}); err != nil { + t.Fatal(err) + } + + if _, err := m.PutIntegrationResponse(ctx(), apiID, resID, "GET", "201", driver.PutIntegrationResponseInput{ + SelectionPattern: "201", + }); err != nil { + t.Fatal(err) + } + + if _, err := m.CreateDeployment(ctx(), apiID, driver.CreateDeploymentInput{StageName: "s"}); err != nil { + t.Fatal(err) + } + + resp := invokeMock(t, m, apiID, driver.ProxyRequest{Body: `{"items":[{"id":1},{"id":2}]}`}) + if resp.StatusCode != 201 { + t.Fatalf("status = %d %s", resp.StatusCode, resp.Body) + } +} + +func TestMappingTemplateSizeQuota(t *testing.T) { + m := newMock(t) + apiID, resID := mockMethod(t, m, map[string]string{"application/json": `{"statusCode": 200}`}, "", "") + big := strings.Repeat("x", 300<<10+1) + + _, err := m.UpdateIntegration(ctx(), apiID, resID, "GET", []driver.PatchOperation{ + {Op: "replace", Path: "/requestTemplates/application~1json", Value: big}, + }) + assertMessage(t, err, errors.IsInvalidArgument, + "Mapping template for content type application/json exceeds the maximum size of 300 KB") + + _, err = m.UpdateIntegrationResponse(ctx(), apiID, resID, "GET", "200", []driver.PatchOperation{ + {Op: "add", Path: "/responseTemplates/text~1plain", Value: big}, + }) + assertMessage(t, err, errors.IsInvalidArgument, + "Mapping template for content type text/plain exceeds the maximum size of 300 KB") +} diff --git a/providers/aws/apigateway/mapping_validation.go b/providers/aws/apigateway/mapping_validation.go index 616fbdb98..862c5e15c 100644 --- a/providers/aws/apigateway/mapping_validation.go +++ b/providers/aws/apigateway/mapping_validation.go @@ -6,6 +6,7 @@ import ( "strings" cerrors "github.com/stackshy/cloudemu/v2/errors" + "github.com/stackshy/cloudemu/v2/internal/vtl" "github.com/stackshy/cloudemu/v2/services/apigateway/driver" ) @@ -27,6 +28,7 @@ const ( mappingErrPrefix = "Invalid mapping expression specified: Validation Result: warnings : [], errors : [" msgBadThroughBehavior = "Invalid passthrough behavior specified" msgInvalidSelection = "Invalid selection pattern specified" + msgTemplateTooLarge = "Mapping template for content type %s exceeds the maximum size of 300 KB" msgContentHandlingEnum = "1 validation error detected: Value '%s' at 'contentHandling' failed to satisfy " + "constraint: Member must satisfy enum value set: [CONVERT_TO_BINARY, CONVERT_TO_TEXT]" ) @@ -84,6 +86,10 @@ func validateIntegrationSettings(ig *driver.Integration, methodParams map[string return err } + if err := validateTemplateSizes(ig.RequestTemplates); err != nil { + return err + } + for _, k := range sortedKeys(ig.RequestParameters) { if !integrationRequestKey.MatchString(k) { return invalidParameter(k) @@ -112,6 +118,10 @@ func validateIntegrationResponse(ir *driver.IntegrationResponse, mr *driver.Meth return cerrors.New(cerrors.InvalidArgument, msgInvalidSelection) } + if err := validateTemplateSizes(ir.ResponseTemplates); err != nil { + return err + } + for _, k := range sortedKeys(ir.ResponseParameters) { declared := mr != nil && mr.ResponseParameters != nil if declared { @@ -131,6 +141,17 @@ func validateIntegrationResponse(ir *driver.IntegrationResponse, mr *driver.Meth return nil } +// validateTemplateSizes enforces API Gateway's 300 KB mapping-template quota. +func validateTemplateSizes(templates map[string]string) error { + for _, ct := range sortedKeys(templates) { + if len(templates[ct]) > vtl.MaxTemplateBytes { + return cerrors.Newf(cerrors.InvalidArgument, msgTemplateTooLarge, ct) + } + } + + return nil +} + func sortedBoolKeys(m map[string]bool) []string { out := make([]string, 0, len(m)) for k := range m { diff --git a/providers/aws/apigateway/mock_integration.go b/providers/aws/apigateway/mock_integration.go index 0203298d3..c7c231f7d 100644 --- a/providers/aws/apigateway/mock_integration.go +++ b/providers/aws/apigateway/mock_integration.go @@ -31,6 +31,7 @@ const ( // body. func (m *Mock) serveMock(ctx context.Context, req *driver.ProxyRequest, route *resolvedRoute) *driver.ProxyResponse { mc := newMappingContext(req, route, m.opts.AccountID, idgen.UUID(), m.opts.Clock.Now()) + mc.templates = m.templates payload, templated, rejected := mapRequest(ctx, mc, &route.integration, req) if rejected != nil { @@ -309,7 +310,9 @@ func bodyPath(body, path string) (string, bool) { switch v.(type) { case *vtl.Map, *vtl.List: - return vtl.ToJSON(v), true + s, err := vtl.ToJSON(v) + + return s, err == nil default: return vtl.Stringify(v), true } From a99f364c9268c3e4951e7f7a3e676e84173484f7 Mon Sep 17 00:00:00 2001 From: Nitin Kumar Date: Sun, 4 Oct 2026 18:54:56 +0530 Subject: [PATCH 3/3] fix(aws-apigateway): budget regex matches, bound concurrent renders, deadline JSONPath walks, byte-capped template cache --- internal/jsonpath/jsonpath.go | 45 +- internal/jsonpath/jsonpath_test.go | 25 +- internal/vtl/cache.go | 70 ++- internal/vtl/cache_test.go | 46 +- internal/vtl/eval.go | 27 +- internal/vtl/limits.go | 46 +- internal/vtl/limits_test.go | 69 +++ internal/vtl/methods.go | 226 +-------- internal/vtl/race_off_test.go | 7 + internal/vtl/race_on_test.go | 7 + internal/vtl/slots_test.go | 64 +++ internal/vtl/string_methods.go | 454 ++++++++++++++++++ providers/aws/apigateway/apigateway.go | 6 +- providers/aws/apigateway/mapping.go | 6 +- .../aws/apigateway/mapping_limits_test.go | 21 + 15 files changed, 876 insertions(+), 243 deletions(-) create mode 100644 internal/vtl/race_off_test.go create mode 100644 internal/vtl/race_on_test.go create mode 100644 internal/vtl/slots_test.go create mode 100644 internal/vtl/string_methods.go diff --git a/internal/jsonpath/jsonpath.go b/internal/jsonpath/jsonpath.go index c2cfc5a80..72d3f14a9 100644 --- a/internal/jsonpath/jsonpath.go +++ b/internal/jsonpath/jsonpath.go @@ -15,6 +15,7 @@ package jsonpath import ( + "context" "fmt" "sort" "strconv" @@ -41,6 +42,26 @@ const maxDepth = 1000 // wildcards and descents cannot multiply a document into a huge result. const maxMatches = 1 << 20 +// ctxCheckEvery is how many visited nodes pass between deadline checks. +const ctxCheckEvery = 1024 + +// walker carries EvalAll's deadline and counts the nodes it visits. +type walker struct { + ctx context.Context + visits int +} + +// visit counts one node and reports the context's error every ctxCheckEvery +// nodes. +func (w *walker) visit() error { + w.visits++ + if w.visits%ctxCheckEvery == 0 { + return w.ctx.Err() + } + + return nil +} + // Error reports a malformed or unsupported path. type Error struct { Msg string @@ -85,8 +106,9 @@ func Eval(path string, root any) (value any, present bool, err error) { // EvalAll evaluates a path that may use wildcards or recursive descent. // indefinite reports whether the path can match more than one value (it uses // a wildcard or descent), in which case callers present the matches as a list. -// For a definite path values holds at most one element. -func EvalAll(path string, root any) (values []any, indefinite bool, err error) { +// For a definite path values holds at most one element. The walk stops with +// ctx's error once ctx is done. +func EvalAll(ctx context.Context, path string, root any) (values []any, indefinite bool, err error) { if rootErr := checkRoot(path); rootErr != nil { return nil, false, rootErr } @@ -101,6 +123,7 @@ func EvalAll(path string, root any) (values []any, indefinite bool, err error) { } cur := []any{root} + w := &walker{ctx: ctx} for _, t := range toks { indefinite = indefinite || t.wildcard || t.descent @@ -108,7 +131,9 @@ func EvalAll(path string, root any) (values []any, indefinite bool, err error) { var next []any for _, v := range cur { - next = t.collect(v, next) + if next, err = t.collect(w, v, next); err != nil { + return nil, true, err + } if len(next) > maxMatches { return nil, true, errorf("JSONPath %q matches more than %d values", path, maxMatches) @@ -149,9 +174,9 @@ func (t token) apply(cur any) (any, bool) { } // collect appends every match of t under v to out. -func (t token) collect(v any, out []any) []any { +func (t token) collect(w *walker, v any, out []any) ([]any, error) { if !t.descent { - return t.collectHere(v, out) + return t.collectHere(v, out), w.visit() } // Walk v and all of its descendants in document order. @@ -168,6 +193,14 @@ func (t token) collect(v any, out []any) []any { out = t.collectHere(f.v, out) + if err := w.visit(); err != nil { + return nil, err + } + + if len(out) > maxMatches { + return out, nil + } + if f.depth >= maxDepth { continue } @@ -178,7 +211,7 @@ func (t token) collect(v any, out []any) []any { } } - return out + return out, nil } // collectHere appends the matches of t directly under v. diff --git a/internal/jsonpath/jsonpath_test.go b/internal/jsonpath/jsonpath_test.go index 4bb3cf749..a146a9a3e 100644 --- a/internal/jsonpath/jsonpath_test.go +++ b/internal/jsonpath/jsonpath_test.go @@ -1,6 +1,7 @@ package jsonpath import ( + "context" "fmt" "sort" "testing" @@ -99,14 +100,14 @@ func TestEvalAll(t *testing.T) { } for _, c := range cases { - got, indefinite, err := EvalAll(c.path, root) + got, indefinite, err := EvalAll(context.Background(), c.path, root) if err != nil || fmt.Sprint(got) != c.want && !(len(got) == 0 && c.want == "[]") || indefinite != c.indefinite { t.Errorf("EvalAll(%q) = %v %v %v, want %s %v", c.path, got, indefinite, err, c.want, c.indefinite) } } for _, bad := range []string{"x", "$[?(@.a)]", "$..", "$.a[", "$[x]"} { - if _, _, err := EvalAll(bad, root); err == nil { + if _, _, err := EvalAll(context.Background(), bad, root); err == nil { t.Errorf("EvalAll(%q) accepted", bad) } } @@ -117,7 +118,7 @@ func TestEvalAll(t *testing.T) { deep = []any{deep} } - if _, _, err := EvalAll("$..*", deep); err != nil { + if _, _, err := EvalAll(context.Background(), "$..*", deep); err != nil { t.Fatalf("deep descent: %v", err) } } @@ -129,7 +130,23 @@ func TestEvalAllCapsMatches(t *testing.T) { items = []any{items, 2} } - if _, _, err := EvalAll("$..*..*..*", items); err == nil { + if _, _, err := EvalAll(context.Background(), "$..*..*..*", items); err == nil { t.Fatal("multiplying path not capped") } } + +func TestEvalAllStopsAtDeadline(t *testing.T) { + // A wide, deep document: descent over it is slow enough to pass a + // cancelled context's first check. + var doc any = 1 + for range 500 { + doc = []any{doc, 1, 2, 3} + } + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + if _, _, err := EvalAll(ctx, "$..*..zz", doc); err == nil { + t.Fatal("cancelled walk did not stop") + } +} diff --git a/internal/vtl/cache.go b/internal/vtl/cache.go index 55ef15b6a..8e380bdb8 100644 --- a/internal/vtl/cache.go +++ b/internal/vtl/cache.go @@ -5,29 +5,44 @@ import ( "sync" ) -// Cache is a bounded, least-recently-used cache of parsed templates keyed by -// their source, so a template is parsed once rather than on every render. A +// Cache weights. A parsed template holds roughly astBytesPerSourceByte bytes +// of syntax tree per byte of source, and every entry costs entryOverhead. +const ( + astBytesPerSourceByte = 20 + entryOverhead = 1 << 10 +) + +// Cache is a least-recently-used cache of parsed templates keyed by their +// source, so a template is parsed once rather than on every render. It is +// bounded by the estimated memory of what it holds, not by entry count. A // parse failure is cached too. A parsed Template is read-only, so one cached // template may render concurrently. type Cache struct { - mu sync.Mutex - max int - order *list.List - items map[string]*list.Element + mu sync.Mutex + maxBytes int + bytes int + order *list.List + items map[string]*list.Element } type cacheEntry struct { - src string - tmpl *Template - err error + src string + tmpl *Template + err error + weight int } -// NewCache returns a cache holding at most maxEntries templates. -func NewCache(maxEntries int) *Cache { - return &Cache{max: maxEntries, order: list.New(), items: map[string]*list.Element{}} +// NewCache returns a cache holding at most about maxBytes of parsed +// templates. +func NewCache(maxBytes int) *Cache { + return &Cache{maxBytes: maxBytes, order: list.New(), items: map[string]*list.Element{}} } -// Parse returns the parsed template for src, parsing it on a miss. +// weight estimates the memory a parsed template of src holds. +func weight(src string) int { return len(src)*astBytesPerSourceByte + entryOverhead } + +// Parse returns the parsed template for src, parsing it on a miss. A template +// heavier than the whole cache is parsed but not kept. func (c *Cache) Parse(src string) (*Template, error) { c.mu.Lock() @@ -42,20 +57,25 @@ func (c *Cache) Parse(src string) (*Template, error) { c.mu.Unlock() tmpl, err := Parse(src) + w := weight(src) c.mu.Lock() defer c.mu.Unlock() - if _, ok := c.items[src]; !ok { - c.items[src] = c.order.PushFront(&cacheEntry{src: src, tmpl: tmpl, err: err}) + if _, ok := c.items[src]; ok || w > c.maxBytes { + return tmpl, err + } + + c.items[src] = c.order.PushFront(&cacheEntry{src: src, tmpl: tmpl, err: err, weight: w}) + c.bytes += w - for c.order.Len() > c.max { - oldest := c.order.Back() - c.order.Remove(oldest) + for c.bytes > c.maxBytes { + oldest := c.order.Back() + c.order.Remove(oldest) - e, _ := oldest.Value.(*cacheEntry) - delete(c.items, e.src) - } + e, _ := oldest.Value.(*cacheEntry) + delete(c.items, e.src) + c.bytes -= e.weight } return tmpl, err @@ -68,3 +88,11 @@ func (c *Cache) Len() int { return c.order.Len() } + +// Bytes returns the estimated memory the cached templates hold. +func (c *Cache) Bytes() int { + c.mu.Lock() + defer c.mu.Unlock() + + return c.bytes +} diff --git a/internal/vtl/cache_test.go b/internal/vtl/cache_test.go index 01cae9c61..b78d12623 100644 --- a/internal/vtl/cache_test.go +++ b/internal/vtl/cache_test.go @@ -2,12 +2,15 @@ package vtl import ( "context" + "runtime" + "strings" "sync" "testing" ) func TestCache(t *testing.T) { - c := NewCache(2) + // Room for two small templates. + c := NewCache(2*entryOverhead + 200) a1, err := c.Parse("a$x") if err != nil { @@ -53,3 +56,44 @@ func TestCache(t *testing.T) { wg.Wait() } + +// TestCacheWeightCoversParsedSize checks the per-byte weight against what a +// parsed template really holds, and that the byte cap bounds the cache. +func TestCacheWeightCoversParsedSize(t *testing.T) { + var b strings.Builder + for b.Len() < MaxTemplateBytes-200 { + b.WriteString(`#if($a.b("x", 1) == [1, 2])$c.d#{else}text $e#end`) + } + + src := b.String() + + var before, after runtime.MemStats + + runtime.GC() + runtime.ReadMemStats(&before) + + tmpl, err := Parse(src) + if err != nil { + t.Fatal(err) + } + + runtime.GC() + runtime.ReadMemStats(&after) + runtime.KeepAlive(tmpl) + + held := int(after.HeapAlloc) - int(before.HeapAlloc) + if held > weight(src) { + t.Fatalf("parsed template holds %d bytes, weight is only %d", held, weight(src)) + } + + c := NewCache(4 * weight(src)) + for i := range 10 { + if _, err := c.Parse(src + strings.Repeat(" ", i)); err != nil { + t.Fatal(err) + } + } + + if c.Len() < 3 || c.Bytes() > 4*weight(src) { + t.Fatalf("cache holds %d entries, %d bytes", c.Len(), c.Bytes()) + } +} diff --git a/internal/vtl/eval.go b/internal/vtl/eval.go index 9a0d76214..f82ecfd48 100644 --- a/internal/vtl/eval.go +++ b/internal/vtl/eval.go @@ -4,6 +4,7 @@ import ( "context" "errors" "math" + "runtime" "strings" ) @@ -18,6 +19,11 @@ const ( ctxCheckEvery = 64 ) +// renderSlots bounds concurrent renders to the CPU count. A host Object must +// not render another template from inside a render, or it could wait on a +// slot its own caller holds. +var renderSlots = make(chan struct{}, runtime.GOMAXPROCS(0)) //nolint:gochecknoglobals // process-wide limit + // ErrStepBudget is returned when a render exceeds its step budget. var ErrStepBudget = errors.New("vtl: template exceeded its execution step budget") @@ -44,7 +50,19 @@ func (t *Template) Render(ctx context.Context, vars map[string]any, opts RenderO vars = map[string]any{} } - st := newState(ctx, vars, opts.MaxSteps, &budget{}) + // Bound how many renders run at once, so concurrent requests cannot + // multiply the per-render memory budget. + select { + case renderSlots <- struct{}{}: + defer func() { <-renderSlots }() + case <-ctx.Done(): + return nil, ctx.Err() + } + + mem := newSharedBudget() + defer mem.release() + + st := newState(ctx, vars, opts.MaxSteps, mem) if st.maxSteps <= 0 { st.maxSteps = DefaultMaxSteps } @@ -59,6 +77,13 @@ func (t *Template) Render(ctx context.Context, vars map[string]any, opts RenderO case errors.As(err, &ret): return &Result{Output: st.out.String(), Returned: true, ReturnValue: ret.value}, nil default: + if errors.Is(err, ErrMemoryLimit) || errors.Is(err, ErrOutputLimit) { + // A render that ran out of budget leaves up to a budget's worth of + // garbage. Collect it now, before concurrent abusive renders pile + // it up faster than the pacer would. + runtime.GC() + } + return nil, err } } diff --git a/internal/vtl/limits.go b/internal/vtl/limits.go index db1cb1713..2ee03b0cb 100644 --- a/internal/vtl/limits.go +++ b/internal/vtl/limits.go @@ -3,6 +3,7 @@ package vtl import ( "errors" "strings" + "sync/atomic" ) // Size and depth limits. A template that hits one fails with an error instead @@ -17,6 +18,9 @@ const ( // MaxAllocBytes caps the strings and collection entries one render may // create, so a loop that keeps doubling a value fails fast. MaxAllocBytes = 64 << 20 + // MaxInFlightAllocBytes caps MaxAllocBytes-style charges summed over every + // render running at once. + MaxInFlightAllocBytes = 2 * MaxAllocBytes // MaxTemplateDepth caps how deeply directives, expressions and string // interpolations may nest. MaxTemplateDepth = 100 @@ -27,8 +31,9 @@ const ( // walks recursively. maxOperatorChain = 1000 // slotBytes is what one list or map entry is charged against - // MaxAllocBytes: the 16-byte interface plus room for slice growth. - slotBytes = 32 + // MaxAllocBytes: the 16-byte interface, the boxed value behind it and room + // for slice growth. + slotBytes = 64 ) // Limit errors. @@ -39,14 +44,45 @@ var ( ErrCyclicValue = errors.New("vtl: cannot encode a value that contains itself") ) -// budget tracks the bytes a render has created. +// inFlightBytes is what all renders running now have charged. It is capped +// at MaxInFlightAllocBytes, so many concurrent renders cannot each use a full +// budget at once. +var inFlightBytes atomic.Int64 //nolint:gochecknoglobals // process-wide memory accounting + +// budget tracks the bytes a render has created. A shared budget also counts +// toward inFlightBytes and must be released when the render ends. type budget struct { - used int + used int + shared bool +} + +func newSharedBudget() *budget { return &budget{shared: true} } + +// release returns a shared budget's bytes to the process-wide pool. +func (b *budget) release() { + if b.shared { + inFlightBytes.Add(-int64(b.used)) + } +} + +// check reports whether n more bytes would fit, without charging them. +func (b *budget) check(n int) error { + if b.used+n > MaxAllocBytes || b.shared && inFlightBytes.Load()+int64(n) > MaxInFlightAllocBytes { + return ErrMemoryLimit + } + + return nil } func (b *budget) charge(n int) error { b.used += n - if b.used > MaxAllocBytes { + + inFlight := int64(0) + if b.shared { + inFlight = inFlightBytes.Add(int64(n)) + } + + if b.used > MaxAllocBytes || inFlight > MaxInFlightAllocBytes { return ErrMemoryLimit } diff --git a/internal/vtl/limits_test.go b/internal/vtl/limits_test.go index 80535d8a9..792725ca2 100644 --- a/internal/vtl/limits_test.go +++ b/internal/vtl/limits_test.go @@ -3,6 +3,9 @@ package vtl import ( "context" "errors" + "fmt" + "regexp" + "runtime" "strings" "testing" "time" @@ -162,3 +165,69 @@ func isLimit(err error) bool { return errors.Is(err, ErrMemoryLimit) || errors.Is(err, ErrOutputLimit) || errors.Is(err, ErrStepBudget) || errors.Is(err, context.DeadlineExceeded) } + +// TestStringMethodsAllocateWithinBudget runs splits and regex replacements over +// a 6 MB string with millions of matches. Each must stop at the memory budget +// and allocate no more than about twice the budget on the way. +func TestStringMethodsAllocateWithinBudget(t *testing.T) { + body := strings.Repeat("&a", 3<<20) + + for name, src := range map[string]string{ + "literal split": `#set($p = $b.split("&"))`, + "regex split": `#set($p = $b.split("[&]"))`, + "regex replace": `#set($p = $b.replaceAll("(&)", "$1"))`, + "empty split": `#set($p = $b.split(""))`, + "anchored split": `#set($p = $b.split("\b"))`, + } { + t.Run(name, func(t *testing.T) { + var before, after runtime.MemStats + + runtime.GC() + runtime.ReadMemStats(&before) + + err := renderErr(t, src, map[string]any{"b": body}) + + runtime.ReadMemStats(&after) + + if !isLimit(err) { + t.Fatalf("err = %v, want a limit error", err) + } + + if alloc := after.TotalAlloc - before.TotalAlloc; !raceEnabled && alloc > 2*MaxAllocBytes { + t.Fatalf("allocated %d MB, budget is %d MB", alloc>>20, MaxAllocBytes>>20) + } + }) + } + + // A modest split still works and keeps Java's trailing-empty rule. + if got := render(t, `$b.split("&").size()`, map[string]any{"b": "a&b&&"}); got != "2" { + t.Fatalf("split size = %s", got) + } + + if got := render(t, `$b.split("")`, map[string]any{"b": "abc"}); got != "[a, b, c]" { + t.Fatalf("empty split = %s", got) + } +} + +// TestEachMatchAgreesWithFindAll checks the incremental matcher against Go's +// FindAll on patterns with empty and overlapping candidates. +func TestEachMatchAgreesWithFindAll(t *testing.T) { + for _, c := range []struct{ pattern, s string }{ + {"a*", "baaacaa"}, {"", "héllo"}, {"x?", "axxbx"}, {"(a)(b)?", "abaab"}, + {"[&]", "&a&&b&"}, {"\\b", "ab cd"}, {"^a", "aaa"}, {"é|", "aéb"}, + } { + re := regexp.MustCompile(c.pattern) + want := fmt.Sprint(re.FindAllStringSubmatchIndex(c.s, -1)) + + var got [][]int + + err := eachMatch(&budget{}, re, c.s, -1, func(loc []int) error { + got = append(got, loc) + + return nil + }) + if err != nil || fmt.Sprint(got) != want { + t.Errorf("%q on %q: got %v, want %s", c.pattern, c.s, got, want) + } + } +} diff --git a/internal/vtl/methods.go b/internal/vtl/methods.go index 552f46af2..95d709434 100644 --- a/internal/vtl/methods.go +++ b/internal/vtl/methods.go @@ -1,10 +1,5 @@ package vtl -import ( - "regexp" - "strings" -) - // Method names shared by more than one receiver type. const ( mIsEmpty = "isEmpty" @@ -18,48 +13,28 @@ const ( mKeySet = "keySet" mValues = "values" mEntrySet = "entrySet" + mAddAll = "addAll" + mPutAll = "putAll" ) // pairArgs is the argument count of a two-argument method such as put or set. const pairArgs = 2 type ( - stringFn func(s string, args []any) (any, error) - listFn func(l *List, args []any) any - mapFn func(m *Map, args []any) any + listFn func(l *List, args []any) any + mapFn func(m *Map, args []any) any ) -// The Java-like method bridge for strings, lists and maps. +// The Java-like method bridge for lists and maps. // //nolint:gochecknoglobals // read-only dispatch tables var ( - stringMethods = map[string]stringFn{ - "length": func(s string, _ []any) (any, error) { return int64(len([]rune(s))), nil }, - mIsEmpty: func(s string, _ []any) (any, error) { return s == "", nil }, - mContains: strPredicate(strings.Contains), - "startsWith": strPredicate(strings.HasPrefix), - "endsWith": strPredicate(strings.HasSuffix), - "equalsIgnoreCase": strPredicate(strings.EqualFold), - mIndexOf: strIndex(strings.Index), - "lastIndexOf": strIndex(strings.LastIndex), - "substring": func(s string, args []any) (any, error) { return substring(s, args), nil }, - "replace": strReplace, - "replaceAll": func(s string, args []any) (any, error) { return regexReplace(s, true, args) }, - "replaceFirst": func(s string, args []any) (any, error) { return regexReplace(s, false, args) }, - "split": strSplit, - "toLowerCase": func(s string, _ []any) (any, error) { return strings.ToLower(s), nil }, - "toUpperCase": func(s string, _ []any) (any, error) { return strings.ToUpper(s), nil }, - "trim": func(s string, _ []any) (any, error) { return strings.TrimSpace(s), nil }, - "matches": strMatches, - "charAt": strCharAt, - } - listMethods = map[string]listFn{ mSize: func(l *List, _ []any) any { return int64(len(l.Items)) }, mIsEmpty: func(l *List, _ []any) any { return len(l.Items) == 0 }, mGet: listGet, "add": listAdd, - "addAll": listAddAll, + mAddAll: listAddAll, mContains: func(l *List, args []any) any { return len(args) == 1 && listIndexOf(l, args[0]) >= 0 }, mIndexOf: listIndexOfMethod, mRemove: listRemove, @@ -69,7 +44,7 @@ var ( mapMethods = map[string]mapFn{ mGet: mapGet, mPut: mapPut, - "putAll": mapPutAll, + mPutAll: mapPutAll, "containsKey": mapContainsKey, mRemove: func(m *Map, args []any) any { return m.Remove(Stringify(firstArg(args))) }, mKeySet: func(m *Map, _ []any) any { return stringList(m.keys) }, @@ -119,7 +94,7 @@ func callString(mem *budget, s, name string, args []any) (any, error) { return nil, nil } - r, err := fn(s, args) + r, err := fn(mem, s, args) if err != nil { return nil, err } @@ -131,8 +106,6 @@ func callString(mem *budget, s, name string, args []any) (any, error) { } return t, mem.charge(len(t)) - case *List: - return t, mem.charge(len(t.Items) * slotBytes) default: return r, nil } @@ -146,6 +119,11 @@ func callCollection(mem *budget, v any, name string, args []any) (any, error) { before = collectionLen(v) ) + // Refuse growth that would not fit before the method allocates it. + if err := mem.check(expectedGrowth(v, name, args) * slotBytes); err != nil { + return nil, err + } + switch t := v.(type) { case *List: fn, ok := listMethods[name] @@ -179,6 +157,19 @@ func callCollection(mem *budget, v any, name string, args []any) (any, error) { return r, nil } +// expectedGrowth estimates the entries a list or map method is about to add +// or build. +func expectedGrowth(v any, name string, args []any) int { + switch { + case name == mAddAll || name == mPutAll: + return collectionLen(firstArg(args)) + case creatingMethods[name]: + return collectionLen(v) + default: + return 1 + } +} + func collectionLen(v any) int { switch t := v.(type) { case *List: @@ -230,171 +221,6 @@ func firstArg(args []any) any { return args[0] } -func strPredicate(pred func(s, arg string) bool) stringFn { - return func(s string, args []any) (any, error) { - a, ok := strArg(args, 0) - - return ok && pred(s, a), nil - } -} - -// strIndex converts a byte offset to a character offset (-1 stays -1). -func strIndex(find func(s, sub string) int) stringFn { - return func(s string, args []any) (any, error) { - a, _ := strArg(args, 0) - - i := find(s, a) - if i < 0 { - return int64(-1), nil - } - - return int64(len([]rune(s[:i]))), nil - } -} - -func strReplace(s string, args []any) (any, error) { - from, ok1 := strArg(args, 0) - to, ok2 := strArg(args, 1) - - if !ok1 || !ok2 { - return nil, nil - } - - return literalReplace(s, from, to, true) -} - -// literalReplace replaces from with to (every occurrence, or the first), -// checking the result size before building it. -func literalReplace(s, from, to string, all bool) (any, error) { - n := 1 - if all { - n = strings.Count(s, from) - } - - if strings.Contains(s, from) && len(s)+n*(len(to)-len(from)) > MaxOutputBytes { - return nil, ErrOutputLimit - } - - if !all { - return strings.Replace(s, from, to, 1), nil - } - - return strings.ReplaceAll(s, from, to), nil -} - -func strMatches(s string, args []any) (any, error) { - pattern, _ := strArg(args, 0) - - re, err := compile("^(?:" + pattern + ")$") - if err != nil { - return nil, err - } - - return re.MatchString(s), nil -} - -func strCharAt(s string, args []any) (any, error) { - i, ok := intArg(args, 0) - r := []rune(s) - - if !ok || i < 0 || i >= len(r) { - return nil, nil - } - - return string(r[i]), nil -} - -func substring(s string, args []any) any { - r := []rune(s) - - begin, ok := intArg(args, 0) - if !ok || begin < 0 || begin > len(r) { - return nil - } - - end := len(r) - - if e, ok := intArg(args, 1); ok { - if e < begin || e > len(r) { - return nil - } - - end = e - } - - return string(r[begin:end]) -} - -func compile(pattern string) (*regexp.Regexp, error) { - re, err := regexp.Compile(pattern) - if err != nil { - return nil, errorf("invalid regular expression %q: %v", pattern, err) - } - - return re, nil -} - -func regexReplace(s string, all bool, args []any) (any, error) { - pattern, _ := strArg(args, 0) - repl, _ := strArg(args, 1) - - // A literal pattern and replacement need no regex engine. - if regexp.QuoteMeta(pattern) == pattern && !strings.ContainsAny(repl, `$\`) && pattern != "" { - return literalReplace(s, pattern, repl, all) - } - - re, err := compile(pattern) - if err != nil { - return nil, err - } - - limit := 1 - if all { - limit = -1 - } - - // Build the result match by match so it can stop at the size limit. - var ( - out []byte - last int - ) - - for _, loc := range re.FindAllStringSubmatchIndex(s, limit) { - out = append(out, s[last:loc[0]]...) - out = re.ExpandString(out, repl, s, loc) - last = loc[1] - - if len(out) > MaxOutputBytes { - return nil, ErrOutputLimit - } - } - - out = append(out, s[last:]...) - if len(out) > MaxOutputBytes { - return nil, ErrOutputLimit - } - - return string(out), nil -} - -// strSplit follows Java's String.split: the argument is a regex and trailing -// empty strings are dropped. -func strSplit(s string, args []any) (any, error) { - pattern, _ := strArg(args, 0) - - re, err := compile(pattern) - if err != nil { - return nil, err - } - - parts := re.Split(s, -1) - for len(parts) > 0 && parts[len(parts)-1] == "" { - parts = parts[:len(parts)-1] - } - - return stringList(parts), nil -} - func listGet(l *List, args []any) any { v, _ := l.Index(intOr(args, -1)) diff --git a/internal/vtl/race_off_test.go b/internal/vtl/race_off_test.go new file mode 100644 index 000000000..636556315 --- /dev/null +++ b/internal/vtl/race_off_test.go @@ -0,0 +1,7 @@ +//go:build !race + +package vtl + +// raceEnabled reports a -race build, where sync.Pool drops items at random and +// allocation counts are not meaningful. +const raceEnabled = false diff --git a/internal/vtl/race_on_test.go b/internal/vtl/race_on_test.go new file mode 100644 index 000000000..6752d1200 --- /dev/null +++ b/internal/vtl/race_on_test.go @@ -0,0 +1,7 @@ +//go:build race + +package vtl + +// raceEnabled reports a -race build, where sync.Pool drops items at random and +// allocation counts are not meaningful. +const raceEnabled = true diff --git a/internal/vtl/slots_test.go b/internal/vtl/slots_test.go new file mode 100644 index 000000000..002a1059f --- /dev/null +++ b/internal/vtl/slots_test.go @@ -0,0 +1,64 @@ +package vtl + +import ( + "context" + "errors" + "testing" + "time" +) + +// TestRenderSlotsBoundConcurrency fills every render slot and checks a further +// render waits for one, giving up when its context ends. +func TestRenderSlotsBoundConcurrency(t *testing.T) { + for range cap(renderSlots) { + renderSlots <- struct{}{} + } + + defer func() { + for range cap(renderSlots) { + <-renderSlots + } + }() + + tmpl, err := Parse("x") + if err != nil { + t.Fatal(err) + } + + ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) + defer cancel() + + if _, err := tmpl.Render(ctx, nil, RenderOptions{}); !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("render with no free slot: err = %v", err) + } +} + +// TestInFlightBudgetIsShared checks that concurrent renders draw on one +// process-wide pool and return their bytes when they end. +func TestInFlightBudgetIsShared(t *testing.T) { + a, b, c := newSharedBudget(), newSharedBudget(), newSharedBudget() + + if err := a.charge(MaxAllocBytes - 1); err != nil { + t.Fatal(err) + } + + if err := b.charge(MaxAllocBytes - 1); err != nil { + t.Fatal(err) + } + + if err := c.check(MaxAllocBytes / 2); !errors.Is(err, ErrMemoryLimit) { + t.Fatalf("third render fit a full pool: %v", err) + } + + if err := c.charge(MaxAllocBytes / 2); !errors.Is(err, ErrMemoryLimit) { + t.Fatalf("third render charged a full pool: %v", err) + } + + a.release() + b.release() + c.release() + + if got := inFlightBytes.Load(); got != 0 { + t.Fatalf("pool holds %d bytes after release", got) + } +} diff --git a/internal/vtl/string_methods.go b/internal/vtl/string_methods.go new file mode 100644 index 000000000..d42c554ab --- /dev/null +++ b/internal/vtl/string_methods.go @@ -0,0 +1,454 @@ +package vtl + +import ( + "regexp" + "regexp/syntax" + "strings" + "unicode/utf8" +) + +// maxPatternBytes caps a regular expression a template passes to a string +// method; compiling a huge pattern costs memory before any match runs. +const maxPatternBytes = 64 << 10 + +// firstMatchBatch is the first batch eachMatchBatched collects; later +// batches double, each checked against the budget before it runs. +const firstMatchBatch = 64 + +// stringFn is one String method. It charges to mem whatever it allocates in +// proportion to its input before allocating it. +type stringFn func(mem *budget, s string, args []any) (any, error) + +// stringMethods is the Java-like String method bridge. +// +//nolint:gochecknoglobals // read-only dispatch table +var stringMethods = map[string]stringFn{ + "length": func(_ *budget, s string, _ []any) (any, error) { return int64(utf8.RuneCountInString(s)), nil }, + mIsEmpty: func(_ *budget, s string, _ []any) (any, error) { return s == "", nil }, + mContains: strPredicate(strings.Contains), + "startsWith": strPredicate(strings.HasPrefix), + "endsWith": strPredicate(strings.HasSuffix), + "equalsIgnoreCase": strPredicate(strings.EqualFold), + mIndexOf: strIndex(strings.Index), + "lastIndexOf": strIndex(strings.LastIndex), + "substring": func(_ *budget, s string, args []any) (any, error) { return substring(s, args), nil }, + "replace": strReplace, + "replaceAll": func(mem *budget, s string, args []any) (any, error) { return regexReplace(mem, s, true, args) }, + "replaceFirst": func(mem *budget, s string, args []any) (any, error) { return regexReplace(mem, s, false, args) }, + "split": strSplit, + "toLowerCase": func(_ *budget, s string, _ []any) (any, error) { return strings.ToLower(s), nil }, + "toUpperCase": func(_ *budget, s string, _ []any) (any, error) { return strings.ToUpper(s), nil }, + "trim": func(_ *budget, s string, _ []any) (any, error) { return strings.TrimSpace(s), nil }, + "matches": strMatches, + "charAt": strCharAt, +} + +func strPredicate(pred func(s, arg string) bool) stringFn { + return func(_ *budget, s string, args []any) (any, error) { + a, ok := strArg(args, 0) + + return ok && pred(s, a), nil + } +} + +// strIndex converts a byte offset to a character offset (-1 stays -1). +func strIndex(find func(s, sub string) int) stringFn { + return func(_ *budget, s string, args []any) (any, error) { + a, _ := strArg(args, 0) + + i := find(s, a) + if i < 0 { + return int64(-1), nil + } + + return int64(utf8.RuneCountInString(s[:i])), nil + } +} + +func strReplace(_ *budget, s string, args []any) (any, error) { + from, ok1 := strArg(args, 0) + to, ok2 := strArg(args, 1) + + if !ok1 || !ok2 { + return nil, nil + } + + return literalReplace(s, from, to, true) +} + +// literalReplace replaces from with to (every occurrence, or the first), +// checking the result size before building it. +func literalReplace(s, from, to string, all bool) (any, error) { + n := 1 + if all { + n = strings.Count(s, from) + } + + if strings.Contains(s, from) && len(s)+n*(len(to)-len(from)) > MaxOutputBytes { + return nil, ErrOutputLimit + } + + if !all { + return strings.Replace(s, from, to, 1), nil + } + + return strings.ReplaceAll(s, from, to), nil +} + +func strMatches(_ *budget, s string, args []any) (any, error) { + pattern, _ := strArg(args, 0) + + re, err := compile("^(?:" + pattern + ")$") + if err != nil { + return nil, err + } + + return re.MatchString(s), nil +} + +// runeAt returns the byte offset of the i-th character of s, or -1 when s has +// fewer than i characters. i == character count yields len(s). +func runeAt(s string, i int) int { + if i < 0 { + return -1 + } + + for off := range s { + if i == 0 { + return off + } + + i-- + } + + if i == 0 { + return len(s) + } + + return -1 +} + +func strCharAt(_ *budget, s string, args []any) (any, error) { + i, ok := intArg(args, 0) + if !ok { + return nil, nil + } + + off := runeAt(s, i) + if off < 0 || off == len(s) { + return nil, nil + } + + _, size := utf8.DecodeRuneInString(s[off:]) + + return s[off : off+size], nil +} + +func substring(s string, args []any) any { + begin, ok := intArg(args, 0) + if !ok { + return nil + } + + from := runeAt(s, begin) + if from < 0 { + return nil + } + + end, hasEnd := intArg(args, 1) + if !hasEnd { + return s[from:] + } + + if end < begin { + return nil + } + + to := runeAt(s, end) + if to < 0 { + return nil + } + + return s[from:to] +} + +func compile(pattern string) (*regexp.Regexp, error) { + if len(pattern) > maxPatternBytes { + return nil, errorf("regular expression is longer than %d bytes", maxPatternBytes) + } + + re, err := regexp.Compile(pattern) + if err != nil { + return nil, errorf("invalid regular expression %q: %v", pattern, err) + } + + return re, nil +} + +// matchCost is what one regex match is charged: its index slice plus the +// regexp engine's per-call state. +const matchCost = 4 * slotBytes + +// eachMatch calls fn for each match of re in s, in order, stopping after n +// matches when n >= 0. It finds one match at a time and charges each before +// the next search, so millions of matches fail at the budget instead of being +// collected up front. It follows FindAll's rules: matches do not overlap, and +// an empty match right after the previous match is skipped. +func eachMatch(mem *budget, re *regexp.Regexp, s string, n int, fn func(loc []int) error) error { + if needsContext(re) { + return eachMatchBatched(mem, re, s, n, fn) + } + + return eachMatchIncremental(mem, re, s, n, fn) +} + +// eachMatchIncremental searches successive suffixes of s, one match at a time. +func eachMatchIncremental(mem *budget, re *regexp.Regexp, s string, n int, fn func(loc []int) error) error { + pos, prevEnd := 0, -1 + + for count := 0; pos <= len(s) && (n < 0 || count < n); { + if err := mem.charge(matchCost); err != nil { + return err + } + + loc := findFrom(re, s, pos) + if loc == nil { + return nil + } + + // FindAll skips an empty match right after the previous match. + if loc[0] != loc[1] || loc[0] != prevEnd { + if err := fn(loc); err != nil { + return err + } + + count++ + prevEnd = loc[1] + } + + var done bool + if pos, done = nextSearch(s, loc); done { + return nil + } + } + + return nil +} + +// findFrom finds the first match of re in s at or after pos, with indexes +// into s. +func findFrom(re *regexp.Regexp, s string, pos int) []int { + loc := re.FindStringSubmatchIndex(s[pos:]) + + for k := range loc { + if loc[k] >= 0 { + loc[k] += pos + } + } + + return loc +} + +// nextSearch returns where to search after match loc: its end, or one +// character further for an empty match. done is true at the end of s. +func nextSearch(s string, loc []int) (pos int, done bool) { + if loc[0] != loc[1] { + return loc[1], false + } + + if loc[1] >= len(s) { + return 0, true + } + + _, w := utf8.DecodeRuneInString(s[loc[1]:]) + + return loc[1] + w, false +} + +// needsContext reports whether re uses an assertion (^, \A, \b, \B) whose +// result depends on the text before the search start, which rules out +// searching a suffix of the string. +func needsContext(re *regexp.Regexp) bool { + parsed, err := syntax.Parse(re.String(), syntax.Perl) + if err != nil { + return true + } + + return hasContextOp(parsed) +} + +func hasContextOp(r *syntax.Regexp) bool { + if r.Op == syntax.OpBeginLine || r.Op == syntax.OpBeginText || + r.Op == syntax.OpWordBoundary || r.Op == syntax.OpNoWordBoundary { + return true + } + + for _, sub := range r.Sub { + if hasContextOp(sub) { + return true + } + } + + return false +} + +// eachMatchBatched serves patterns with context assertions: it collects +// matches with FindAll in doubling batches, checking each batch against the +// budget before running it. +func eachMatchBatched(mem *budget, re *regexp.Regexp, s string, n int, fn func(loc []int) error) error { + for batch := firstMatchBatch; ; batch *= 2 { + if n >= 0 && batch >= n { + batch = n + } + + if err := mem.check(2 * batch * matchCost); err != nil { + return err + } + + locs := re.FindAllStringSubmatchIndex(s, batch) + if len(locs) == batch && batch != n { + continue + } + + if err := mem.charge(len(locs) * matchCost); err != nil { + return err + } + + for _, loc := range locs { + if err := fn(loc); err != nil { + return err + } + } + + return nil + } +} + +func regexReplace(mem *budget, s string, all bool, args []any) (any, error) { + pattern, _ := strArg(args, 0) + repl, _ := strArg(args, 1) + + // A literal pattern and replacement need no regex engine. + if regexp.QuoteMeta(pattern) == pattern && !strings.ContainsAny(repl, `$\`) && pattern != "" { + return literalReplace(s, pattern, repl, all) + } + + re, err := compile(pattern) + if err != nil { + return nil, err + } + + n := 1 + if all { + n = -1 + } + + // Build the result match by match so it can stop at the size limit. + var ( + out []byte + last int + ) + + err = eachMatch(mem, re, s, n, func(loc []int) error { + out = append(out, s[last:loc[0]]...) + out = re.ExpandString(out, repl, s, loc) + last = loc[1] + + if len(out) > MaxOutputBytes { + return ErrOutputLimit + } + + return nil + }) + if err != nil { + return nil, err + } + + out = append(out, s[last:]...) + if len(out) > MaxOutputBytes { + return nil, ErrOutputLimit + } + + return string(out), nil +} + +// strSplit follows Java's String.split: the argument is a regex and trailing +// empty strings are dropped. Each piece is charged as it is added, so a huge +// split stops at the budget. +func strSplit(mem *budget, s string, args []any) (any, error) { + pattern, _ := strArg(args, 0) + + var ( + l *List + err error + ) + + if regexp.QuoteMeta(pattern) == pattern && pattern != "" { + l, err = splitLiteral(mem, s, pattern) + } else { + l, err = splitRegex(mem, s, pattern) + } + + if err != nil { + return nil, err + } + + for len(l.Items) > 0 && l.Items[len(l.Items)-1] == "" { + l.Items = l.Items[:len(l.Items)-1] + } + + return l, nil +} + +func splitLiteral(mem *budget, s, sep string) (*List, error) { + l := NewList() + + for { + if err := mem.charge(slotBytes); err != nil { + return nil, err + } + + i := strings.Index(s, sep) + if i < 0 { + l.Items = append(l.Items, s) + + return l, nil + } + + l.Items = append(l.Items, s[:i]) + s = s[i+len(sep):] + } +} + +func splitRegex(mem *budget, s, pattern string) (*List, error) { + re, err := compile(pattern) + if err != nil { + return nil, err + } + + l := NewList() + last := 0 + + err = eachMatch(mem, re, s, -1, func(loc []int) error { + // Java drops the empty leading piece a zero-width match at 0 makes. + if loc[1] == 0 { + return nil + } + + if cerr := mem.charge(slotBytes); cerr != nil { + return cerr + } + + l.Items = append(l.Items, s[last:loc[0]]) + last = loc[1] + + return nil + }) + if err != nil { + return nil, err + } + + l.Items = append(l.Items, s[last:]) + + return l, mem.charge(slotBytes) +} diff --git a/providers/aws/apigateway/apigateway.go b/providers/aws/apigateway/apigateway.go index cdeaab8f1..263bce832 100644 --- a/providers/aws/apigateway/apigateway.go +++ b/providers/aws/apigateway/apigateway.go @@ -90,15 +90,15 @@ type Mock struct { templates *vtl.Cache } -// templateCacheSize bounds the parsed mapping templates kept in memory. -const templateCacheSize = 512 +// templateCacheBytes bounds the memory of the parsed mapping templates kept. +const templateCacheBytes = 64 << 20 // New creates a new API Gateway mock. func New(opts *config.Options) *Mock { return &Mock{ apis: memstore.New[*apiData](), opts: opts, certs: map[string]*driver.ClientCertificate{}, account: defaultAccount(), - templates: vtl.NewCache(templateCacheSize), + templates: vtl.NewCache(templateCacheBytes), } } diff --git a/providers/aws/apigateway/mapping.go b/providers/aws/apigateway/mapping.go index 5820fa056..e954b4cbc 100644 --- a/providers/aws/apigateway/mapping.go +++ b/providers/aws/apigateway/mapping.go @@ -102,7 +102,7 @@ func (mc *mappingContext) render(ctx context.Context, src, body string) (string, } vars := map[string]any{ - "input": &inputObject{body: body, params: mc.params()}, + "input": &inputObject{ctx: ctx, body: body, params: mc.params()}, "context": mc.context, "stageVariables": vtl.StringMap(mc.route.stageVariables), "util": utilObject{}, @@ -161,6 +161,8 @@ func (mc *mappingContext) responseOverride() (status int, headers map[string]str // inputObject is $input. type inputObject struct { + // ctx is the render's deadline, which bounds JSONPath walks. + ctx context.Context body string params *vtl.Map parsed any @@ -219,7 +221,7 @@ func (in *inputObject) path(p string) (any, error) { } } - matches, indefinite, err := jsonpath.EvalAll(p, in.parsed) + matches, indefinite, err := jsonpath.EvalAll(in.ctx, p, in.parsed) if err != nil { return nil, err } diff --git a/providers/aws/apigateway/mapping_limits_test.go b/providers/aws/apigateway/mapping_limits_test.go index eb618d5fb..b8e6345ae 100644 --- a/providers/aws/apigateway/mapping_limits_test.go +++ b/providers/aws/apigateway/mapping_limits_test.go @@ -3,6 +3,7 @@ package apigateway_test import ( "strings" "testing" + "time" "github.com/stackshy/cloudemu/v2/errors" "github.com/stackshy/cloudemu/v2/services/apigateway/driver" @@ -100,3 +101,23 @@ func TestMappingTemplateSizeQuota(t *testing.T) { assertMessage(t, err, errors.IsInvalidArgument, "Mapping template for content type text/plain exceeds the maximum size of 300 KB") } + +// TestMockJSONPathWalkHonoursDeadline runs a descent that multiplies over a +// deep, wide body. The template deadline must end it as a 500. +func TestMockJSONPathWalkHonoursDeadline(t *testing.T) { + m := newMock(t) + reqTmpl := `#set($x = $input.path('$..*..*..zz')){"statusCode": 200}` + apiID, _ := mockMethod(t, m, map[string]string{"application/json": reqTmpl}, "", "") + + body := strings.Repeat("[1,2,3,", 900) + "0" + strings.Repeat("]", 900) + start := time.Now() + + resp := invokeMock(t, m, apiID, driver.ProxyRequest{Body: body}) + if resp.StatusCode != 500 { + t.Fatalf("status = %d %s", resp.StatusCode, resp.Body) + } + + if d := time.Since(start); d > 4*time.Second { + t.Fatalf("walk took %v", d) + } +}