package cabana
import (
"context"
"encoding/json"
"errors"
"html/template"
"net/http"
"net/http/httptest"
"os"
"reflect"
"strings"
"testing"
"testing/fstest"
"time"
"git.golem15.com/golem15/summercms/modules/bouncer"
"git.golem15.com/golem15/summercms/modules/pact"
"git.golem15.com/golem15/summercms/modules/towel"
"golang.org/x/net/html"
"golang.org/x/net/html/atom"
"gorm.io/gorm"
)
// sanitizeHTML runs raw markup through the same fragment parse and allowlist
// walk a rendered partial goes through, without html/template in front (which
// would already strip comments).
func sanitizeHTML(t *testing.T, src string) []PartialNode {
t.Helper()
container := &html.Node{Type: html.ElementNode, Data: "div", DataAtom: atom.Div}
parsed, err := html.ParseFragment(strings.NewReader(src), container)
if err != nil {
t.Fatalf("parse %q: %v", src, err)
}
nodes, err := sanitizePartialNodes(parsed, 0, &partialBudget{nodes: partialMaxNodes})
if err != nil {
t.Fatalf("sanitize %q: %v", src, err)
}
return nodes
}
func nodesJSON(t *testing.T, nodes []PartialNode) string {
t.Helper()
raw, err := json.Marshal(nodes)
if err != nil {
t.Fatal(err)
}
return string(raw)
}
// renderPartial parses src as a partial template and renders it once.
func renderPartial(t *testing.T, src string, data any) ([]PartialNode, error) {
t.Helper()
compiled, err := parsePartial("probe", []byte(src))
if err != nil {
t.Fatalf("parse template: %v", err)
}
return compiled.render(context.Background(), nil, data)
}
// TestPhase101PartialSanitizer covers the server half of T-10.1-08, T-10.1-09
// and T-10.1-12: the tag, attribute and URL allowlist, escaping of view-model
// data, per-request translation, the size, node and depth caps, and the
// view-model guard.
func TestPhase101PartialSanitizer(t *testing.T) {
t.Run("dropped tags go with their subtree", func(t *testing.T) {
for tag := range partialDroppedTags {
var src string
switch tag {
case "input", "link", "meta", "base", "embed":
src = `
keep
<` + tag + ` value="gone" href="/gone" content="gone">tail
`
case "svg", "math":
src = `keep
<` + tag + `>gone gone ` + tag + `>tail
`
case "select":
src = `keep
gone tail
`
default:
src = `keep
<` + tag + `>gone gone ` + tag + `>tail
`
}
got := nodesJSON(t, sanitizeHTML(t, src))
if strings.Contains(got, "gone") || strings.Contains(got, `"`+tag+`"`) {
t.Fatalf("%s survived: %s", tag, got)
}
if !strings.Contains(got, `"text":"keep"`) || !strings.Contains(got, `"text":"tail"`) {
t.Fatalf("%s took its siblings with it: %s", tag, got)
}
}
if len(partialDroppedTags) != 19 {
t.Fatalf("dropped tag list has %d entries, the test expects 19", len(partialDroppedTags))
}
})
t.Run("unknown elements are unwrapped, children kept", func(t *testing.T) {
nodes := sanitizeHTML(t, `kept text `)
want := []PartialNode{{Tag: "span", Attrs: map[string]string{"class": "x"}, Children: []PartialNode{{Text: "kept"}}}, {Text: " text"}}
if !reflect.DeepEqual(nodes, want) {
t.Fatalf("nodes = %s", nodesJSON(t, nodes))
}
})
t.Run("attribute allowlist", func(t *testing.T) {
nodes := sanitizeHTML(t, `v
`)
want := map[string]string{"class": "c", "title": "t", "lang": "pl", "dir": "rtl", "role": "note", "aria-label": "a", "data-count": "3"}
if len(nodes) != 1 || !reflect.DeepEqual(nodes[0].Attrs, want) {
t.Fatalf("attrs = %s", nodesJSON(t, nodes))
}
cells := nodesJSON(t, sanitizeHTML(t, `d m p n `))
for _, want := range []string{`"colspan":"2"`, `"rowspan":"1"`, `"scope":"row"`, `"datetime":"2026-09-29"`, `"optimum":"1"`, `"low":"0"`, `"high":"2"`, `"max":"3"`, `"value":"7"`} {
if !strings.Contains(cells, want) {
t.Fatalf("%s missing from %s", want, cells)
}
}
if strings.Contains(cells, "width") || strings.Contains(cells, "onload") {
t.Fatalf("per-tag attributes leaked: %s", cells)
}
// A per-tag attribute is not global: colspan on a div is dropped.
if got := sanitizeHTML(t, `v
`); got[0].Attrs != nil {
t.Fatalf("div attrs = %v", got[0].Attrs)
}
})
t.Run("link and image URLs", func(t *testing.T) {
for _, tc := range []struct {
src, attr string
keep bool
}{
{`l `, "href", true},
{`l `, "href", true},
{`l `, "href", true},
{`l `, "href", false},
{`l `, "href", false},
{`l `, "href", false},
{`l `, "href", false},
{`l `, "href", false},
{`l `, "href", false},
{"l ", "href", false},
{`l `, "href", false},
{`l `, "href", false},
{` `, "src", true},
{` `, "src", false},
{` `, "src", false},
{` `, "src", false},
} {
nodes := sanitizeHTML(t, tc.src)
if len(nodes) != 1 {
t.Fatalf("%s: nodes = %s", tc.src, nodesJSON(t, nodes))
}
_, kept := nodes[0].Attrs[tc.attr]
if kept != tc.keep {
t.Fatalf("%s: %s kept=%v, want %v (%v)", tc.src, tc.attr, kept, tc.keep, nodes[0].Attrs)
}
}
if got := sanitizeHTML(t, ` `); !reflect.DeepEqual(got[0].Attrs, map[string]string{"src": "/a.png", "alt": "a", "width": "10", "height": "10"}) {
t.Fatalf("img attrs = %v", got[0].Attrs)
}
})
t.Run("comments and doctypes are dropped", func(t *testing.T) {
got := nodesJSON(t, sanitizeHTML(t, `ab
`))
if strings.Contains(got, "secret") || strings.Contains(got, "also") || strings.Contains(got, "html") {
t.Fatalf("comment leaked: %s", got)
}
nodes, err := renderPartial(t, "ab
", nil)
if err != nil || strings.Contains(nodesJSON(t, nodes), "template comment") {
t.Fatalf("rendered comment: %v %s", err, nodesJSON(t, nodes))
}
})
t.Run("view-model markup stays text", func(t *testing.T) {
hostile := `bold `
nodes, err := renderPartial(t, `{{ .Data.Name }}
t `, struct{ Name string }{hostile})
if err != nil {
t.Fatal(err)
}
want := []PartialNode{
{Tag: "p", Attrs: map[string]string{"class": "n"}, Children: []PartialNode{{Text: hostile}}},
{Tag: "span", Attrs: map[string]string{"title": hostile}, Children: []PartialNode{{Text: "t"}}},
}
if !reflect.DeepEqual(nodes, want) {
t.Fatalf("nodes = %s", nodesJSON(t, nodes))
}
})
t.Run("trans resolves per request on one compiled partial", func(t *testing.T) {
_, tr := extTranslator(t)
_, cc := mustCompileExt(t, newExtController(), os.DirFS(extDir))
stats := cc.partials["stats"]
vm, _ := extPartialView("stats", nil)
for _, tc := range []struct{ locale, label string }{{"en", "All gadgets"}, {"pl", "Wszystkie gadżety"}, {"en", "All gadgets"}} {
nodes, err := stats.render(towel.WithLocale(context.Background(), tc.locale), tr, vm)
if err != nil {
t.Fatalf("%s render: %v", tc.locale, err)
}
got := nodesJSON(t, nodes)
if !strings.Contains(got, `"text":"`+tc.label+`"`) || !strings.Contains(got, `"text":"3"`) || !strings.Contains(got, `"class":"summer-stats"`) {
t.Fatalf("%s nodes = %s", tc.locale, got)
}
}
// The pristine template was never executed, so it can still be cloned.
if _, err := stats.pristine.Clone(); err != nil {
t.Fatalf("pristine template was executed: %v", err)
}
})
t.Run("caps", func(t *testing.T) {
big := strings.Repeat("x", partialMaxBytes+1)
if _, err := renderPartial(t, `{{ .Data }}
`, big); !errors.Is(err, errPartialTooLarge) {
t.Fatalf("size cap err = %v", err)
}
if _, err := renderPartial(t, `{{ .Data }}
`, strings.Repeat("x", partialMaxBytes-len("
"))); err != nil {
t.Fatalf("output at the size cap: %v", err)
}
// Every x is two nodes: the element and its text.
atCap := strings.Repeat("x ", partialMaxNodes/2)
if _, err := renderPartial(t, atCap, nil); err != nil {
t.Fatalf("%d nodes: %v", partialMaxNodes, err)
}
if _, err := renderPartial(t, atCap+"y", nil); err == nil || !strings.Contains(err.Error(), "2000 nodes") {
t.Fatalf("node cap err = %v", err)
}
nested := func(depth int) string {
return strings.Repeat("", depth) + "x" + strings.Repeat("
", depth)
}
if _, err := renderPartial(t, nested(partialMaxDepth), nil); err != nil {
t.Fatalf("depth %d: %v", partialMaxDepth, err)
}
if _, err := renderPartial(t, nested(partialMaxDepth+1), nil); err == nil || !strings.Contains(err.Error(), "depth 32") {
t.Fatalf("depth cap err = %v", err)
}
})
t.Run("view-model guard", func(t *testing.T) {
_, cc := mustCompileExt(t, newExtController(), os.DirFS(extDir))
for name, vm := range map[string]any{
"model": extGadget{},
"pointer": &extGadget{},
"slice": []extGadget{{}},
"slice of ptr": []*extGadget{{}},
"map of model": map[string]*extGadget{},
"template.HTML": struct{ Body template.HTML }{},
"nested HTMLAttr": struct {
Items []struct{ A template.HTMLAttr }
}{},
"template.URL": map[string]template.URL{},
// WR-03: the model nested in a wrapper, behind an interface, as
// another GORM model, or reached through a method.
"wrapped model": struct{ Gadget *extGadget }{},
"wrapped model slice": struct{ Rows []struct{ G extGadget } }{},
"model in map[string]any": map[string]any{"gadget": &extGadget{Name: "top-secret"}},
"model in []any": []any{1, extGadget{}},
"model in an any field": struct{ Row any }{Row: []*extGadget{{}}},
"HTML in map[string]any": map[string]any{"banner": template.HTML("x ")},
"another tabler model": struct{ User *BackendUser }{},
"gorm-tagged struct": struct{ Row vmTagged }{},
"embedded gorm.Model": struct{ vmGormModel }{},
"gorm.DeletedAt": struct{ Deleted gorm.DeletedAt }{},
"method returns HTML": vmHTMLMethod{},
"pointer method HTML": struct{ Inner vmPtrHTMLMethod }{},
"method with an argument": vmArgHTMLMethod{},
"method returns the model": vmModelMethod{},
} {
if refusedViewModel(cc, vm) == "" {
t.Fatalf("%s view model was accepted", name)
}
}
selfRef := &vmNode{Label: "a"}
selfRef.Next = selfRef
loop := map[string]any{"n": 1}
loop["self"] = loop
for name, vm := range map[string]any{
"nil": nil,
"curated": struct{ Name string }{},
"items": struct{ Items []extStatItem }{},
"string": "text",
// The shape of a curated statistics view model: labels and
// integers, including behind interfaces.
"stats view": struct {
Total int
Formats []extStatItem
NoShelf int
}{Total: 3, Formats: []extStatItem{{Label: "x", Count: 3}}},
"curated map": map[string]any{"total": 3, "items": []extStatItem{{}}, "label": "x", "none": nil},
"time field": struct{ At time.Time }{At: time.Now()},
"self reference": selfRef,
"cyclic map": loop,
"plain method": vmPlainMethod{},
"nil pointer field": struct{ Next *vmNode }{},
} {
if reason := refusedViewModel(cc, vm); reason != "" {
t.Fatalf("%s view model refused: %s", name, reason)
}
}
})
t.Run("the partial route refuses a model view model", func(t *testing.T) {
app, _ := extTranslator(t)
reg, _ := mustCompileExt(t, leakyController{newExtController()}, os.DirFS(extDir))
svc := &service{app: app, reg: reg}
// _summary.htm reads .Data.Name, which the model has too: without
// the guard the model would render.
req := httptest.NewRequest(http.MethodGet, adminAPI("/acme/demo/gadgets/partials/summary"), nil)
req.SetPathValue("vendor", "acme")
req.SetPathValue("plugin", "demo")
req.SetPathValue("controller", "gadgets")
req.SetPathValue("name", "summary")
req = req.WithContext(bouncer.WithUser(req.Context(), &bouncer.Principal{ID: 1, Backend: true, IsSuperuser: true}))
rec := httptest.NewRecorder()
svc.partial(rec, req)
if rec.Code != http.StatusInternalServerError || strings.Contains(rec.Body.String(), "top-secret") {
t.Fatalf("status=%d body=%s", rec.Code, rec.Body.String())
}
assertErrorCode(t, rec.Body.Bytes(), "error")
})
}
// leakyController hands its GORM model to the template, which the partial
// handler must refuse.
type leakyController struct{ extController }
func (leakyController) PartialData(context.Context, string, any) (any, error) {
return &extGadget{Name: "top-secret"}, nil
}
// View-model fixtures of the guard (WR-03).
type vmTagged struct {
Secret string `gorm:"column:secret"`
}
type vmGormModel struct{ gorm.Model }
type vmHTMLMethod struct{ Body string }
func (v vmHTMLMethod) Banner() template.HTML { return template.HTML(v.Body) }
type vmPtrHTMLMethod struct{ Body string }
func (v *vmPtrHTMLMethod) Banner() template.HTML { return template.HTML(v.Body) }
type vmArgHTMLMethod struct{}
func (vmArgHTMLMethod) Wrap(s string) template.HTML { return template.HTML(s) }
type vmModelMethod struct{}
func (vmModelMethod) Gadget() *extGadget { return &extGadget{} }
type vmPlainMethod struct{ N int }
func (v vmPlainMethod) Double() int { return v.N * 2 }
type vmNode struct {
Label string
Next *vmNode
}
// TestPhase101WidgetContext covers WR-06: the widget route honours the
// field's context like the save path. Without record_id the request is the
// create form's, with one the update form's; a widget its context hides on
// that form answers 404 and its action never runs.
func TestPhase101WidgetContext(t *testing.T) {
const lookup = ` lookup:
label: acme.demo::lang.gadgets.lookup
type: widget
widget: acme-demo-lookup
action: lookup
fill: [name, active]
`
for _, tc := range []struct {
context string
body string
want int
}{
{context: "update", body: `{}`, want: http.StatusNotFound},
{context: "[update, preview]", body: `{"values":{}}`, want: http.StatusNotFound},
{context: "create", body: `{"record_id":1}`, want: http.StatusNotFound},
{context: "create", body: `{}`, want: http.StatusOK},
{context: "", body: `{}`, want: http.StatusOK},
} {
t.Run(tc.context+" "+tc.body, func(t *testing.T) {
fsys := extFS(t)
fields := string(fsys["models/gadget/fields.yaml"].Data)
if !strings.Contains(fields, lookup) {
t.Fatalf("fixture fields.yaml changed:\n%s", fields)
}
if tc.context != "" {
fields = strings.Replace(fields, lookup, lookup+" context: "+tc.context+"\n", 1)
}
fsys["models/gadget/fields.yaml"] = &fstest.MapFile{Data: []byte(fields)}
ran := 0
ctl := newExtController()
ctl.actions = []pact.AdminAction{{
Name: "lookup", Label: "acme.demo::lang.gadgets.lookup",
Run: func(context.Context, pact.AdminActionInput) (pact.AdminActionResult, error) {
ran++
return pact.AdminActionResult{}, nil
},
}, extActions()[1]}
app, _ := extTranslator(t)
reg, _ := mustCompileExt(t, ctl, fsys)
svc := &service{app: app, reg: reg}
req := httptest.NewRequest(http.MethodPost, adminAPI("/acme/demo/gadgets/widgets/lookup"), strings.NewReader(tc.body))
req.SetPathValue("vendor", "acme")
req.SetPathValue("plugin", "demo")
req.SetPathValue("controller", "gadgets")
req.SetPathValue("field", "lookup")
req = req.WithContext(bouncer.WithUser(req.Context(), &bouncer.Principal{ID: 1, Backend: true, IsSuperuser: true}))
rec := httptest.NewRecorder()
svc.widgetAction(rec, req)
if rec.Code != tc.want {
t.Fatalf("status=%d body=%s, want %d", rec.Code, rec.Body.String(), tc.want)
}
if wantRan := tc.want == http.StatusOK; (ran == 1) != wantRan {
t.Fatalf("action ran %d times, want ran=%v", ran, wantRan)
}
if tc.want == http.StatusNotFound {
assertErrorCode(t, rec.Body.Bytes(), "not_found")
}
})
}
}