Files
summercms/tide/variables.go
2026-09-17 14:43:01 +02:00

720 lines
17 KiB
Go

package tide
import (
"encoding/json"
"fmt"
"net/http"
"net/url"
"os"
"path/filepath"
"regexp"
"sort"
"strconv"
"strings"
"sync"
"github.com/goccy/go-yaml"
)
var (
placeholderRe = regexp.MustCompile(`\{\{([^{}]+)\}\}`)
jwtShapeRe = regexp.MustCompile(`eyJ[A-Za-z0-9_-]+\.[A-Za-z0-9_-]+\.[A-Za-z0-9_-]+`)
invShapeRe = regexp.MustCompile(`inv_[A-Za-z0-9]{8,}`)
cookieRe = regexp.MustCompile(`(?i)auth_token=([^;]+)`)
secretFormRe = regexp.MustCompile(`(?i)client_secret=([^&\s]+)`)
pkceFormRe = regexp.MustCompile(`(?i)code_verifier=([^&\s]+)`)
passwordJSONRe = regexp.MustCompile(`(?i)"password"\s*:\s*"([^"]*)"`)
passwordFormRe = regexp.MustCompile(`(?i)(?:^|&)password=([^&\s]*)`)
accessTokenJSONRe = regexp.MustCompile(`(?i)"access_token"\s*:\s*"([^"]*)"`)
secretJSONRe = regexp.MustCompile(`(?i)"client_secret"\s*:\s*"[^"]*"`)
)
// allowedTestPasswords are documented onboarding secrets that remain in fixtures
// as plaintext. Any other leftover password-shaped value is a capture leak.
var allowedTestPasswords = map[string]struct{}{
"parity-alice-pass": {},
}
// Store holds named capture values. When Path is set it is a mode-0600 private file.
type Store struct {
mu sync.Mutex
path string
vals map[string]string
}
// OpenStore loads or creates a private variable map. Empty path is memory-only.
func OpenStore(path string) (*Store, error) {
s := &Store{path: path, vals: make(map[string]string)}
if strings.TrimSpace(path) == "" {
return s, nil
}
abs, err := filepath.Abs(path)
if err != nil {
return nil, fmt.Errorf("tide: vars path: %w", err)
}
s.path = abs
st, err := os.Stat(abs)
if err == nil {
if st.IsDir() {
return nil, fmt.Errorf("tide: vars %q is a directory", path)
}
raw, err := os.ReadFile(abs)
if err != nil {
return nil, fmt.Errorf("tide: read vars: %w", err)
}
if len(strings.TrimSpace(string(raw))) > 0 {
if err := yaml.Unmarshal(raw, &s.vals); err != nil {
return nil, fmt.Errorf("tide: parse vars: %w", err)
}
if s.vals == nil {
s.vals = make(map[string]string)
}
}
if err := os.Chmod(abs, 0o600); err != nil {
return nil, fmt.Errorf("tide: chmod vars: %w", err)
}
return s, nil
}
if !os.IsNotExist(err) {
return nil, fmt.Errorf("tide: stat vars: %w", err)
}
if err := os.MkdirAll(filepath.Dir(abs), 0o755); err != nil {
return nil, fmt.Errorf("tide: create vars dir: %w", err)
}
if err := os.WriteFile(abs, []byte("{}\n"), 0o600); err != nil {
return nil, fmt.Errorf("tide: create vars: %w", err)
}
_ = os.Chmod(abs, 0o600)
return s, nil
}
// Save writes the map as YAML with mode 0600. Memory-only stores are a no-op.
func (s *Store) Save() error {
if s == nil || s.path == "" {
return nil
}
s.mu.Lock()
defer s.mu.Unlock()
keys := make([]string, 0, len(s.vals))
for k := range s.vals {
keys = append(keys, k)
}
sort.Strings(keys)
var b strings.Builder
if len(keys) == 0 {
b.WriteString("{}\n")
}
for _, k := range keys {
fmt.Fprintf(&b, "%s: %s\n", strconv.Quote(k), strconv.Quote(s.vals[k]))
}
raw := []byte(b.String())
tmp, err := os.CreateTemp(filepath.Dir(s.path), ".vars-*.tmp")
if err != nil {
return fmt.Errorf("tide: vars temp: %w", err)
}
tmpName := tmp.Name()
if _, err := tmp.Write(raw); err != nil {
_ = tmp.Close()
_ = os.Remove(tmpName)
return err
}
_ = tmp.Chmod(0o600)
if err := tmp.Close(); err != nil {
_ = os.Remove(tmpName)
return err
}
if err := os.Rename(tmpName, s.path); err != nil {
_ = os.Remove(tmpName)
return err
}
return os.Chmod(s.path, 0o600)
}
// Get returns a stored value.
func (s *Store) Get(name string) (string, bool) {
if s == nil {
return "", false
}
s.mu.Lock()
defer s.mu.Unlock()
v, ok := s.vals[name]
return v, ok
}
// Set stores a named value.
func (s *Store) Set(name, value string) {
if s == nil {
return
}
s.mu.Lock()
defer s.mu.Unlock()
s.vals[name] = value
}
// Expand replaces {{name}} placeholders. Unresolved names fail before HTTP send.
func (s *Store) Expand(text string) (string, error) {
if !strings.Contains(text, "{{") {
return text, nil
}
if s == nil {
m := placeholderRe.FindStringSubmatch(text)
if len(m) > 1 {
return "", fmt.Errorf("tide: unresolved placeholder %q", m[1])
}
return "", fmt.Errorf("tide: unresolved placeholder")
}
var missing []string
s.mu.Lock()
out := placeholderRe.ReplaceAllStringFunc(text, func(m string) string {
name := m[2 : len(m)-2]
v, ok := s.vals[name]
if !ok {
missing = append(missing, name)
return m
}
return v
})
s.mu.Unlock()
if len(missing) > 0 {
return "", fmt.Errorf("tide: unresolved placeholder %q", missing[0])
}
return out, nil
}
func expandRequest(req Request, store *Store) (Request, error) {
out := req
var err error
out.Path, err = store.Expand(req.Path)
if err != nil {
return Request{}, err
}
out.Query, err = store.Expand(req.Query)
if err != nil {
return Request{}, err
}
if req.Headers != nil {
out.Headers = make(map[string]string, len(req.Headers))
for k, v := range req.Headers {
out.Headers[k], err = store.Expand(v)
if err != nil {
return Request{}, err
}
}
}
body, err := store.Expand(string(req.Body))
if err != nil {
return Request{}, err
}
out.Body = Body(body)
return out, nil
}
func expandResponse(resp Response, store *Store) (Response, error) {
out := resp
var err error
if resp.Headers != nil {
out.Headers = make(map[string]string, len(resp.Headers))
for k, v := range resp.Headers {
out.Headers[k], err = store.Expand(v)
if err != nil {
return Response{}, err
}
}
}
body, err := store.Expand(string(resp.Body))
if err != nil {
return Response{}, err
}
out.Body = Body(body)
return out, nil
}
// CaptureStep writes named values from the step into the store.
func CaptureStep(store *Store, step *Step) error {
if store == nil || step == nil || len(step.Capture) == 0 {
return nil
}
for _, rule := range step.Capture {
val, err := extractCapture(rule, step.Request, step.Response)
if err != nil {
return fmt.Errorf("tide: capture %q: %w", rule.As, err)
}
if strings.TrimSpace(val) == "" {
return fmt.Errorf("tide: capture %q was empty", rule.As)
}
store.Set(rule.As, val)
}
return nil
}
func extractCapture(rule CaptureRule, req Request, resp Response) (string, error) {
from := strings.TrimSpace(rule.From)
if from == "" {
from = "response.json"
}
switch from {
case "response.json":
v, err := jsonPathValue([]byte(resp.Body), rule.Path)
if err != nil {
return "", err
}
return scalarString(v), nil
case "response.json.query":
v, err := jsonPathValue([]byte(resp.Body), rule.Path)
if err != nil {
return "", err
}
raw := scalarString(v)
u, err := url.Parse(raw)
if err != nil {
return "", err
}
q := u.Query().Get(rule.Name)
if q == "" {
return "", fmt.Errorf("missing JSON URL query %s", rule.Name)
}
return q, nil
case "response.header":
v := headerValue(resp.Headers, rule.Name)
if v == "" {
return "", fmt.Errorf("missing response header %s", rule.Name)
}
return v, nil
case "response.query":
loc := headerValue(resp.Headers, "Location")
if loc == "" {
return "", fmt.Errorf("missing Location header")
}
u, err := url.Parse(loc)
if err != nil {
return "", err
}
v := u.Query().Get(rule.Name)
if v == "" {
return "", fmt.Errorf("missing response query %s", rule.Name)
}
return v, nil
case "response.location.query":
loc := headerValue(resp.Headers, "Location")
if loc == "" {
return "", fmt.Errorf("missing Location header")
}
u, err := url.Parse(loc)
if err != nil {
return "", err
}
v := u.Query().Get(rule.Name)
if v == "" {
return "", fmt.Errorf("missing Location query %s", rule.Name)
}
return v, nil
case "request.form":
v, err := formValue(string(req.Body), rule.Name)
if err != nil {
return "", err
}
if v == "" {
return "", fmt.Errorf("missing form field %s", rule.Name)
}
return v, nil
case "request.header":
v := headerValue(req.Headers, rule.Name)
if v == "" {
return "", fmt.Errorf("missing request header %s", rule.Name)
}
return v, nil
case "request.query":
q, err := url.ParseQuery(req.Query)
if err != nil {
return "", err
}
v := q.Get(rule.Name)
if v == "" {
return "", fmt.Errorf("missing request query %s", rule.Name)
}
return v, nil
default:
return "", fmt.Errorf("unknown capture from %q", from)
}
}
func formValue(body, name string) (string, error) {
vals, err := url.ParseQuery(body)
if err != nil {
return "", err
}
return vals.Get(name), nil
}
func jsonPathValue(raw []byte, path string) (any, error) {
if strings.TrimSpace(path) == "" {
return nil, fmt.Errorf("json path is required")
}
root, err := decodeJSON(raw)
if err != nil {
return nil, err
}
v, err := walkJSONPath(root, path)
if err != nil {
return nil, err
}
return v, nil
}
func walkJSONPath(root any, path string) (any, error) {
path = strings.TrimSpace(path)
if path == "$" || path == "" {
return root, nil
}
if !strings.HasPrefix(path, "$") {
path = "$." + path
}
cur := root
rest := strings.TrimPrefix(path, "$")
for rest != "" {
switch {
case strings.HasPrefix(rest, "."):
rest = rest[1:]
name, next := splitPathSeg(rest)
if name == "" {
return nil, fmt.Errorf("invalid json path %s", path)
}
obj, ok := cur.(map[string]any)
if !ok {
return nil, fmt.Errorf("%s is not an object", path)
}
v, ok := obj[name]
if !ok {
return nil, fmt.Errorf("missing %s", "$."+name)
}
cur = v
rest = next
case strings.HasPrefix(rest, "["):
end := strings.IndexByte(rest, ']')
if end < 0 {
return nil, fmt.Errorf("invalid json path %s", path)
}
idx, err := atoi(rest[1:end])
if err != nil {
return nil, err
}
arr, ok := cur.([]any)
if !ok || idx < 0 || idx >= len(arr) {
return nil, fmt.Errorf("missing %s[%d]", path, idx)
}
cur = arr[idx]
rest = rest[end+1:]
default:
return nil, fmt.Errorf("invalid json path %s", path)
}
}
return cur, nil
}
func splitPathSeg(s string) (name, rest string) {
i := 0
for i < len(s) && s[i] != '.' && s[i] != '[' {
i++
}
return s[:i], s[i:]
}
func atoi(s string) (int, error) {
n := 0
if s == "" {
return 0, fmt.Errorf("empty index")
}
for _, c := range s {
if c < '0' || c > '9' {
return 0, fmt.Errorf("invalid index %q", s)
}
n = n*10 + int(c-'0')
}
return n, nil
}
func scalarString(v any) string {
switch t := v.(type) {
case nil:
return ""
case string:
return t
case json.Number:
return string(t)
case bool:
return fmt.Sprintf("%v", t)
default:
return fmt.Sprintf("%v", t)
}
}
// ScrubStep replaces stored capture values with {{name}} in kept fields.
func ScrubStep(store *Store, step *Step) error {
if store == nil || step == nil {
return nil
}
pairs := store.replacements()
// Short numeric IDs belong in path/query/headers (e.g. /albums/1). JSON
// bodies keep literal counts and IPv4; normalizeJSON already masks id/*_id.
step.Request.Path = replaceAll(step.Request.Path, pairs, true)
step.Request.Query = replaceAll(step.Request.Query, pairs, true)
step.Request.Headers = scrubMap(step.Request.Headers, pairs, true)
step.Request.Body = Body(replaceAll(string(step.Request.Body), pairs, false))
for _, rule := range step.Capture {
if strings.TrimSpace(rule.From) != "request.form" {
continue
}
if rule.Name == "" || rule.As == "" {
continue
}
step.Request.Body = Body(scrubFormField(string(step.Request.Body), rule.Name, rule.As))
}
step.Response.Headers = scrubMap(step.Response.Headers, pairs, true)
step.Response.Body = Body(replaceAll(string(step.Response.Body), pairs, false))
return rejectUnclassifiedCredentials(*step)
}
func (s *Store) replacements() [][2]string {
if s == nil {
return nil
}
s.mu.Lock()
defer s.mu.Unlock()
keys := make([]string, 0, len(s.vals))
for k, v := range s.vals {
if v != "" {
keys = append(keys, k)
}
}
sort.Slice(keys, func(i, j int) bool {
return len(s.vals[keys[i]]) > len(s.vals[keys[j]])
})
out := make([][2]string, 0, len(keys))
for _, k := range keys {
out = append(out, [2]string{s.vals[k], "{{" + k + "}}"})
}
return out
}
func scrubMap(in map[string]string, pairs [][2]string, allowShortNumeric bool) map[string]string {
if in == nil {
return nil
}
out := make(map[string]string, len(in))
for k, v := range in {
out[k] = replaceAll(v, pairs, allowShortNumeric)
}
return out
}
func replaceAll(s string, pairs [][2]string, allowShortNumeric bool) string {
for _, p := range pairs {
if p[0] == "" {
continue
}
if !allowShortNumeric && isAllDigits(p[0]) && len(p[0]) < 8 {
continue
}
olds := []string{p[0]}
if esc := phpJSONEscape(p[0]); esc != p[0] {
olds = append(olds, esc)
}
for _, old := range olds {
if len(old) >= 8 {
s = strings.ReplaceAll(s, old, p[1])
continue
}
s = replaceIsolated(s, old, p[1])
}
}
return s
}
func isAllDigits(s string) bool {
if s == "" {
return false
}
for i := 0; i < len(s); i++ {
if s[i] < '0' || s[i] > '9' {
return false
}
}
return true
}
func phpJSONEscape(s string) string {
return strings.ReplaceAll(s, "/", `\/`)
}
func scrubFormField(body, name, as string) string {
if name == "" || as == "" || body == "" {
return body
}
re := regexp.MustCompile(`(?i)(^|&)(` + regexp.QuoteMeta(name) + `=)[^&]*`)
return re.ReplaceAllString(body, `${1}${2}{{`+as+`}}`)
}
func replaceIsolated(s, old, neu string) string {
if old == "" || s == "" {
return s
}
var b strings.Builder
i := 0
for i < len(s) {
j := strings.Index(s[i:], old)
if j < 0 {
b.WriteString(s[i:])
break
}
j += i
leftOK := j == 0 || !isIdentByte(s[j-1])
right := j + len(old)
rightOK := right == len(s) || !isIdentByte(s[right])
if leftOK && rightOK {
b.WriteString(s[i:j])
b.WriteString(neu)
i = right
continue
}
b.WriteString(s[i : j+len(old)])
i = j + len(old)
}
return b.String()
}
func isIdentByte(c byte) bool {
return (c >= '0' && c <= '9') || (c >= 'A' && c <= 'Z') ||
(c >= 'a' && c <= 'z') || c == '_' || c == '.'
}
func rejectUnclassifiedCredentials(step Step) error {
check := func(label, s string) error {
if hit := remainingCredential(s); hit != "" {
return fmt.Errorf("unclassified credential-shaped value (%s) in %s step %s", hit, label, step.ID)
}
return nil
}
for k, v := range step.Request.Headers {
if err := check("request header "+k, v); err != nil {
return err
}
}
if err := check("request query", step.Request.Query); err != nil {
return err
}
if err := check("request body", string(step.Request.Body)); err != nil {
return err
}
for k, v := range step.Response.Headers {
if err := check("response header "+k, v); err != nil {
return err
}
}
return check("response body", string(step.Response.Body))
}
func remainingCredential(s string) string {
s = placeholderRe.ReplaceAllString(s, "")
if s == "" {
return ""
}
if jwtShapeRe.MatchString(s) {
return "jwt"
}
if invShapeRe.MatchString(s) {
return "token"
}
if cookieRe.MatchString(s) {
return "cookie"
}
if secretFormRe.MatchString(s) {
return "oauth_secret"
}
if pkceFormRe.MatchString(s) {
return "pkce"
}
if leftoverPassword(s) {
return "password"
}
if leftoverAccessToken(s) {
return "access_token"
}
return ""
}
func leftoverPassword(s string) bool {
for _, m := range passwordJSONRe.FindAllStringSubmatch(s, -1) {
if m[1] != "" {
if _, ok := allowedTestPasswords[m[1]]; !ok {
return true
}
}
}
for _, m := range passwordFormRe.FindAllStringSubmatch(s, -1) {
v, err := url.QueryUnescape(m[1])
if err != nil {
v = m[1]
}
if v != "" {
if _, ok := allowedTestPasswords[v]; !ok {
return true
}
}
}
return false
}
func leftoverAccessToken(s string) bool {
for _, m := range accessTokenJSONRe.FindAllStringSubmatch(s, -1) {
if strings.TrimSpace(m[1]) != "" {
return true
}
}
return false
}
func redactSecrets(s string) string {
s = jwtShapeRe.ReplaceAllString(s, "<redacted-jwt>")
s = invShapeRe.ReplaceAllString(s, "<redacted-inv>")
s = secretFormRe.ReplaceAllString(s, "client_secret=<redacted>")
s = secretJSONRe.ReplaceAllString(s, `"client_secret":"<redacted>"`)
s = cookieRe.ReplaceAllString(s, "auth_token=<redacted>")
return s
}
func varsOutsideFixtures(varsPath, fixtures string) error {
if varsPath == "" || fixtures == "" {
return nil
}
absVars, err := filepath.Abs(varsPath)
if err != nil {
return fmt.Errorf("tide: vars path: %w", err)
}
absFix, err := filepath.Abs(fixtures)
if err != nil {
return fmt.Errorf("tide: fixtures path: %w", err)
}
if absVars == absFix || strings.HasPrefix(absVars, absFix+string(os.PathSeparator)) {
return fmt.Errorf("tide: vars file %q must be outside fixtures %q", varsPath, fixtures)
}
return nil
}
func mergeRouteCaptures(step *Step, rules Rules) {
if step == nil || len(step.Capture) > 0 {
return
}
route := rules.Match(step.Request.Method, step.Request.Path)
if route != nil && len(route.Capture) > 0 {
step.Capture = append([]CaptureRule(nil), route.Capture...)
}
}
func recordedResponseHeaders(h http.Header, rules Rules, method, path string) map[string]string {
route := rules.Match(method, path)
if keep := rules.responseHeaders(route); len(keep) > 0 {
return filterHeaders(h, keep)
}
return keepResponseHeaders(h)
}