Files
summercms/modules/tide/upstream.go
Jakub Zych ee0004fb65 feat(14-01): record vendor calls with summer parity:upstream and replay them offline
- WriteUpstream masks vars, hashes long base64 JSON strings and refuses unmasked Authorization/X-Api-Key
- multipart requests recorded as ordered parts; the fake compares parts and hashed payloads
- loopback CONNECT recording proxy with a local ECDSA parity CA, script and forward modes
- parity:upstream command, README and parity docs
2026-10-03 19:55:42 +02:00

683 lines
21 KiB
Go

package tide
import (
"bytes"
"crypto/sha256"
"encoding/base64"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
"mime"
"mime/multipart"
"net/http"
"net/url"
"os"
"path/filepath"
"regexp"
"slices"
"strings"
"sync"
"github.com/goccy/go-yaml"
)
// UpstreamSidecar is the upstream HTTP traffic a backend sent to outside
// services while one parity fixture was recorded. It lives next to the
// fixture as <fixture>.upstream.yaml (see UpstreamPath). Exchanges are
// consumed in recorded order across the whole flow.
type UpstreamSidecar struct {
Version int `yaml:"version"`
Exchanges []UpstreamExchange `yaml:"exchanges"`
}
// UpstreamExchange is one outbound request and the vendor's answer.
type UpstreamExchange struct {
Request UpstreamRequest `yaml:"request"`
Response UpstreamResponse `yaml:"response"`
}
// UpstreamRequest is the request the backend sent upstream. URL is absolute
// (scheme, host, path and query). Headers holds only the compared headers
// (UpstreamCompareHeaders); credential values are {{name}} placeholders.
// A multipart/form-data request keeps Parts instead of Body.
type UpstreamRequest struct {
Method string `yaml:"method"`
URL string `yaml:"url"`
Headers map[string]string `yaml:"headers,omitempty"`
Body string `yaml:"body,omitempty"`
Parts []UpstreamPart `yaml:"parts,omitempty"`
}
// UpstreamPart is one multipart/form-data part of a recorded request, in
// order. A plain field keeps its Value (credentials masked as {{name}}); a
// file part keeps Filename, ContentType and the SHA256 of its bytes, never
// the bytes.
type UpstreamPart struct {
Name string `yaml:"name"`
Value string `yaml:"value,omitempty"`
Filename string `yaml:"filename,omitempty"`
ContentType string `yaml:"content_type,omitempty"`
SHA256 string `yaml:"sha256,omitempty"`
}
// MaxUpstreamBody caps one upstream request or response body the fake and
// the recording proxy read. Vision requests carry a base64 photo, so it is
// larger than DefaultMaxBody.
const MaxUpstreamBody = 32 << 20
// upstreamHashMin is the length above which a base64 JSON string value is
// stored as a {{sha256:<hex>}} placeholder.
const upstreamHashMin = 1024
// UpstreamResponse is the recorded vendor answer the fake replays.
type UpstreamResponse struct {
Status int `yaml:"status"`
Headers map[string]string `yaml:"headers,omitempty"`
Body string `yaml:"body,omitempty"`
}
// UpstreamCompareHeaders are the request headers the fake asserts. A header
// is compared when either the recorded or the sent request carries it.
var UpstreamCompareHeaders = []string{
"User-Agent",
"Accept",
"Content-Type",
"Authorization",
"X-Api-Key",
"Anthropic-Version",
"Anthropic-Beta",
}
// UpstreamPath maps a fixture path to its sidecar path: x.yaml becomes
// x.upstream.yaml.
func UpstreamPath(fixturePath string) string {
for _, ext := range []string{".yaml", ".yml"} {
if strings.HasSuffix(fixturePath, ext) {
return strings.TrimSuffix(fixturePath, ext) + ".upstream.yaml"
}
}
return fixturePath + ".upstream.yaml"
}
// LoadUpstream reads a version-1 sidecar, rejecting unknown fields. A missing
// file returns an error wrapping fs.ErrNotExist.
func LoadUpstream(path string) (UpstreamSidecar, error) {
raw, err := os.ReadFile(path)
if err != nil {
return UpstreamSidecar{}, fmt.Errorf("tide: read upstream %s: %w", path, err)
}
var s UpstreamSidecar
dec := yaml.NewDecoder(bytes.NewReader(raw), yaml.DisallowUnknownField())
if err := dec.Decode(&s); err != nil {
return UpstreamSidecar{}, fmt.Errorf("tide: parse upstream %s: %w", path, err)
}
if err := validateUpstream(s); err != nil {
return UpstreamSidecar{}, fmt.Errorf("tide: upstream %s: %w", path, err)
}
return s, nil
}
func validateUpstream(s UpstreamSidecar) error {
if s.Version != CurrentVersion {
return fmt.Errorf("version %d, want %d", s.Version, CurrentVersion)
}
for i, ex := range s.Exchanges {
if strings.TrimSpace(ex.Request.Method) == "" {
return fmt.Errorf("exchange %d: request method is required", i)
}
u, err := url.Parse(ex.Request.URL)
if err != nil || u.Scheme == "" || u.Host == "" {
return fmt.Errorf("exchange %d: request url %q must be absolute", i, ex.Request.URL)
}
if ex.Response.Status < 100 || ex.Response.Status > 999 {
return fmt.Errorf("exchange %d: response status %d out of range", i, ex.Response.Status)
}
}
return nil
}
// UpstreamFake is an http.RoundTripper that answers from a sidecar without
// dialing. Each request takes the next recorded exchange and is asserted
// against it: method, scheme, host, path, query (order-insensitive), the
// UpstreamCompareHeaders with {{name}} placeholders expanded from the store,
// and the body (JSON semantically, otherwise byte for byte). A mismatch fails
// the request and is remembered for Verify.
type UpstreamFake struct {
mu sync.Mutex
store *Store
exchanges []UpstreamExchange
next int
errs []error
}
// NewUpstreamFake returns a fake that replays s. store resolves {{name}}
// placeholders in the recorded requests and responses; it may be nil when the
// sidecar has none.
func NewUpstreamFake(s UpstreamSidecar, store *Store) *UpstreamFake {
return &UpstreamFake{store: store, exchanges: slices.Clone(s.Exchanges)}
}
// RoundTrip implements http.RoundTripper.
func (f *UpstreamFake) RoundTrip(req *http.Request) (*http.Response, error) {
var body []byte
if req.Body != nil {
var err error
body, err = io.ReadAll(io.LimitReader(req.Body, MaxUpstreamBody+1))
_ = req.Body.Close()
if err != nil {
return nil, f.fail(fmt.Errorf("tide: upstream %s %s: read body: %w", req.Method, req.URL, err))
}
}
f.mu.Lock()
if f.next >= len(f.exchanges) {
f.mu.Unlock()
return nil, f.fail(fmt.Errorf("tide: upstream extra request %s %s", req.Method, redactURL(req.URL)))
}
idx := f.next
ex := f.exchanges[idx]
f.next++
f.mu.Unlock()
if problems := f.compare(ex.Request, req, body); len(problems) > 0 {
return nil, f.fail(fmt.Errorf("tide: upstream exchange %d %s %s: %s",
idx, req.Method, redactURL(req.URL), strings.Join(problems, "; ")))
}
resp, err := f.response(ex.Response, req)
if err != nil {
return nil, f.fail(fmt.Errorf("tide: upstream exchange %d: %w", idx, err))
}
return resp, nil
}
// Verify reports every mismatch, every extra request and every recorded
// exchange that was never requested. It returns nil when the backend sent
// exactly the recorded requests.
func (f *UpstreamFake) Verify() error {
f.mu.Lock()
defer f.mu.Unlock()
errs := slices.Clone(f.errs)
for i := f.next; i < len(f.exchanges); i++ {
r := f.exchanges[i].Request
errs = append(errs, fmt.Errorf("tide: upstream exchange %d unconsumed: %s %s", i, r.Method, r.URL))
}
return errors.Join(errs...)
}
func (f *UpstreamFake) fail(err error) error {
f.mu.Lock()
f.errs = append(f.errs, err)
f.mu.Unlock()
return err
}
func (f *UpstreamFake) compare(want UpstreamRequest, got *http.Request, body []byte) []string {
var problems []string
if !strings.EqualFold(want.Method, got.Method) {
problems = append(problems, fmt.Sprintf("method: want %s, got %s", want.Method, got.Method))
}
rawURL, err := f.store.Expand(want.URL)
if err != nil {
return append(problems, "url: "+err.Error())
}
wu, err := url.Parse(rawURL)
if err != nil {
return append(problems, "url: "+err.Error())
}
gu := got.URL
if !strings.EqualFold(wu.Scheme, gu.Scheme) {
problems = append(problems, fmt.Sprintf("scheme: want %s, got %s", wu.Scheme, gu.Scheme))
}
if !strings.EqualFold(wu.Host, gu.Host) {
problems = append(problems, fmt.Sprintf("host: want %s, got %s", wu.Host, gu.Host))
}
if wu.EscapedPath() != gu.EscapedPath() && wu.Path != gu.Path {
problems = append(problems, fmt.Sprintf("path: want %s, got %s", wu.EscapedPath(), gu.EscapedPath()))
}
problems = append(problems, compareQuery(wu.Query(), gu.Query())...)
for _, name := range UpstreamCompareHeaders {
wv, err := f.store.Expand(headerValue(want.Headers, name))
if err != nil {
problems = append(problems, "header "+name+": "+err.Error())
continue
}
gv := got.Header.Get(name)
if wv == gv || (strings.EqualFold(name, "Content-Type") && sameMultipart(wv, gv)) {
continue
}
if isCredentialHeader(name) {
problems = append(problems, fmt.Sprintf("header %s: value differs", name))
continue
}
problems = append(problems, fmt.Sprintf("header %s: want %q, got %q", name, wv, gv))
}
if len(want.Parts) > 0 {
return append(problems, f.compareParts(want.Parts, got.Header.Get("Content-Type"), body)...)
}
wantBody, err := f.store.expandKeeping(want.Body, isHashPlaceholder)
if err != nil {
return append(problems, "body: "+err.Error())
}
jsonBody := isJSONContentType(want.Headers) || isJSONContentType(map[string]string{"Content-Type": got.Header.Get("Content-Type")})
switch {
case jsonBody && (len(wantBody) > 0 || len(body) > 0):
if d, ok := diffUpstreamJSON([]byte(wantBody), body); !ok {
problems = append(problems, fmt.Sprintf("body %s: want %s, got %s", d.Path, clip(d.Expected), clip(d.Actual)))
}
case wantBody != string(body):
problems = append(problems, fmt.Sprintf("body: want %d bytes, got %d bytes", len(wantBody), len(body)))
}
return problems
}
func compareQuery(want, got url.Values) []string {
var problems []string
keys := make([]string, 0, len(want)+len(got))
for k := range want {
keys = append(keys, k)
}
for k := range got {
if _, ok := want[k]; !ok {
keys = append(keys, k)
}
}
slices.Sort(keys)
for _, k := range keys {
wv, gv := slices.Clone(want[k]), slices.Clone(got[k])
slices.Sort(wv)
slices.Sort(gv)
if !slices.Equal(wv, gv) {
problems = append(problems, fmt.Sprintf("query %s: want %q, got %q", k, wv, gv))
}
}
return problems
}
func (f *UpstreamFake) response(r UpstreamResponse, req *http.Request) (*http.Response, error) {
h := http.Header{}
for k, v := range r.Headers {
ev, err := f.store.Expand(v)
if err != nil {
return nil, fmt.Errorf("response header %s: %w", k, err)
}
h.Set(k, ev)
}
body, err := f.store.Expand(r.Body)
if err != nil {
return nil, fmt.Errorf("response body: %w", err)
}
return &http.Response{
Status: fmt.Sprintf("%d %s", r.Status, http.StatusText(r.Status)),
StatusCode: r.Status,
Proto: "HTTP/1.1",
ProtoMajor: 1,
ProtoMinor: 1,
Header: h,
Body: io.NopCloser(strings.NewReader(body)),
ContentLength: int64(len(body)),
Request: req,
}, nil
}
func isCredentialHeader(name string) bool {
return strings.EqualFold(name, "Authorization") || strings.EqualFold(name, "X-Api-Key")
}
// redactURL drops the query, which may carry a vendor token.
func redactURL(u *url.URL) string {
if u == nil {
return ""
}
c := *u
c.RawQuery = ""
c.User = nil
return c.String()
}
// sameMultipart reports whether both values are multipart/form-data, whose
// boundary differs on every request.
func sameMultipart(a, b string) bool {
return isMultipart(a) && isMultipart(b)
}
func isMultipart(contentType string) bool {
media, _, err := mime.ParseMediaType(contentType)
return err == nil && media == "multipart/form-data"
}
func (f *UpstreamFake) compareParts(want []UpstreamPart, contentType string, body []byte) []string {
got, err := upstreamParts(contentType, body)
if err != nil {
return []string{"parts: " + err.Error()}
}
var problems []string
if len(want) != len(got) {
problems = append(problems, fmt.Sprintf("parts: want %d, got %d", len(want), len(got)))
}
for i := range min(len(want), len(got)) {
w, g := want[i], got[i]
value, err := f.store.Expand(w.Value)
if err != nil {
problems = append(problems, fmt.Sprintf("part %d %s: %v", i, w.Name, err))
continue
}
switch {
case w.Name != g.Name:
problems = append(problems, fmt.Sprintf("part %d name: want %q, got %q", i, w.Name, g.Name))
case w.Filename != g.Filename:
problems = append(problems, fmt.Sprintf("part %d %s filename: want %q, got %q", i, w.Name, w.Filename, g.Filename))
case w.ContentType != g.ContentType:
problems = append(problems, fmt.Sprintf("part %d %s content type: want %q, got %q", i, w.Name, w.ContentType, g.ContentType))
case !strings.EqualFold(w.SHA256, g.SHA256):
problems = append(problems, fmt.Sprintf("part %d %s sha256: want %s, got %s", i, w.Name, w.SHA256, g.SHA256))
case value != g.Value:
problems = append(problems, fmt.Sprintf("part %d %s value differs", i, w.Name))
}
}
return problems
}
// upstreamParts splits a multipart/form-data body into recorded parts: plain
// fields keep their value, file parts their sha256.
func upstreamParts(contentType string, body []byte) ([]UpstreamPart, error) {
media, params, err := mime.ParseMediaType(contentType)
if err != nil || media != "multipart/form-data" || params["boundary"] == "" {
return nil, fmt.Errorf("content type %q is not multipart/form-data", contentType)
}
mr := multipart.NewReader(bytes.NewReader(body), params["boundary"])
var parts []UpstreamPart
for {
p, err := mr.NextRawPart()
if errors.Is(err, io.EOF) {
return parts, nil
}
if err != nil {
return nil, err
}
raw, err := io.ReadAll(p)
if err != nil {
return nil, err
}
up := UpstreamPart{Name: p.FormName(), Filename: p.FileName()}
if up.Filename != "" || p.Header.Get("Content-Type") != "" {
up.ContentType = p.Header.Get("Content-Type")
sum := sha256.Sum256(raw)
up.SHA256 = hex.EncodeToString(sum[:])
} else {
up.Value = string(raw)
}
parts = append(parts, up)
}
}
var hashPlaceholderRe = regexp.MustCompile(`^(data:[^,]*;base64,)?\{\{sha256:([0-9a-f]{64})\}\}$`)
func isHashPlaceholder(name string) bool { return strings.HasPrefix(name, "sha256:") }
// expandKeeping is Expand that leaves placeholders whose name keep accepts.
func (s *Store) expandKeeping(text string, keep func(string) bool) (string, error) {
if !strings.Contains(text, "{{") {
return text, nil
}
var missing []string
out := placeholderRe.ReplaceAllStringFunc(text, func(m string) string {
name := m[2 : len(m)-2]
if keep(name) {
return m
}
v, ok := s.Get(name)
if !ok {
missing = append(missing, name)
return m
}
return v
})
if len(missing) > 0 {
return "", fmt.Errorf("tide: unresolved placeholder %q", missing[0])
}
return out, nil
}
// diffUpstreamJSON compares JSON semantically. A recorded string of the form
// {{sha256:<hex>}} (optionally behind a data: URL prefix) matches a sent
// base64 string whose decoded bytes hash to <hex>.
func diffUpstreamJSON(want, got []byte) (Diff, bool) {
wv, err := decodeJSON(want)
if err != nil {
return Diff{Path: "$", Expected: "valid JSON", Actual: err.Error()}, false
}
gv, err := decodeJSON(got)
if err != nil {
return Diff{Path: "$", Expected: formatValue(wv), Actual: err.Error()}, false
}
gv = resolveHashes(wv, gv)
var diffs []Diff
compareValue("$", wv, gv, &diffs)
if len(diffs) > 0 {
return diffs[0], false
}
return Diff{}, true
}
func resolveHashes(want, got any) any {
switch w := want.(type) {
case map[string]any:
g, ok := got.(map[string]any)
if !ok {
return got
}
out := make(map[string]any, len(g))
for k, v := range g {
if wv, ok := w[k]; ok {
v = resolveHashes(wv, v)
}
out[k] = v
}
return out
case []any:
g, ok := got.([]any)
if !ok {
return got
}
out := slices.Clone(g)
for i := range min(len(w), len(g)) {
out[i] = resolveHashes(w[i], g[i])
}
return out
case string:
m := hashPlaceholderRe.FindStringSubmatch(w)
g, ok := got.(string)
if m == nil || !ok || !strings.HasPrefix(g, m[1]) {
return got
}
raw, ok := decodeBase64(strings.TrimPrefix(g, m[1]))
if !ok {
return got
}
sum := sha256.Sum256(raw)
if hex.EncodeToString(sum[:]) == m[2] {
return w
}
return got
}
return got
}
func decodeBase64(s string) ([]byte, bool) {
for _, enc := range []*base64.Encoding{base64.StdEncoding, base64.RawStdEncoding, base64.URLEncoding, base64.RawURLEncoding} {
if raw, err := enc.DecodeString(s); err == nil {
return raw, true
}
}
return nil, false
}
func clip(s string) string {
const max = 120
if len(s) <= max {
return s
}
return s[:max] + fmt.Sprintf("... (%d bytes)", len(s))
}
// WriteUpstream masks s and writes it to path (mode 0644). Every header value,
// URL and body substring equal to a store variable value becomes {{name}};
// a JSON string longer than 1024 characters that decodes as base64 (also
// behind a data: URL prefix) becomes {{sha256:<hex of the decoded bytes>}}.
// An Authorization or X-Api-Key request header that is not fully masked
// refuses the write, naming the header but never its value.
func WriteUpstream(path string, s UpstreamSidecar, store *Store) error {
if s.Version == 0 {
s.Version = CurrentVersion
}
masked, err := maskUpstream(s, store)
if err != nil {
return err
}
if err := validateUpstream(masked); err != nil {
return fmt.Errorf("tide: upstream %s: %w", path, err)
}
raw, err := yaml.MarshalWithOptions(masked, yaml.UseLiteralStyleIfMultiline(true))
if err != nil {
return fmt.Errorf("tide: marshal upstream: %w", err)
}
dir := filepath.Dir(path)
if err := os.MkdirAll(dir, 0o755); err != nil {
return fmt.Errorf("tide: create upstream dir: %w", err)
}
tmp, err := os.CreateTemp(dir, ".upstream-*.tmp")
if err != nil {
return fmt.Errorf("tide: upstream temp: %w", err)
}
tmpName := tmp.Name()
if _, err := tmp.Write(raw); err != nil {
_ = tmp.Close()
_ = os.Remove(tmpName)
return fmt.Errorf("tide: write upstream: %w", err)
}
if err := tmp.Chmod(0o644); err != nil {
_ = tmp.Close()
_ = os.Remove(tmpName)
return err
}
if err := tmp.Close(); err != nil {
_ = os.Remove(tmpName)
return err
}
if err := os.Rename(tmpName, path); err != nil {
_ = os.Remove(tmpName)
return fmt.Errorf("tide: write upstream: %w", err)
}
return nil
}
func maskUpstream(s UpstreamSidecar, store *Store) (UpstreamSidecar, error) {
pairs := store.replacements()
out := UpstreamSidecar{Version: s.Version, Exchanges: make([]UpstreamExchange, len(s.Exchanges))}
for i, ex := range s.Exchanges {
req := ex.Request
req.URL = replaceAll(req.URL, pairs, true)
req.Headers = scrubMap(req.Headers, pairs, true)
if isJSONContentType(req.Headers) {
req.Body = hashBase64Strings(req.Body)
}
req.Body = replaceAll(req.Body, pairs, false)
if len(req.Parts) > 0 {
req.Parts = slices.Clone(req.Parts)
for j := range req.Parts {
req.Parts[j].Value = replaceAll(req.Parts[j].Value, pairs, false)
}
}
for name, v := range req.Headers {
if isCredentialHeader(name) && !fullyMasked(v) {
return UpstreamSidecar{}, fmt.Errorf("tide: upstream exchange %d: %s header is not masked by a vars entry; refusing to write a live credential", i, name)
}
}
resp := ex.Response
resp.Headers = scrubMap(resp.Headers, pairs, true)
resp.Body = replaceAll(resp.Body, pairs, false)
out.Exchanges[i] = UpstreamExchange{Request: req, Response: resp}
}
return out, nil
}
// credentialResidueRe is what may remain of a credential header once its
// placeholders are removed: nothing, or a scheme word optionally followed by
// one key= label (for example "Bearer" or "Discogs token=").
var credentialResidueRe = regexp.MustCompile(`^(?:[A-Za-z]+(?:\s+[A-Za-z_]+=)?)?$`)
func fullyMasked(v string) bool {
v = strings.TrimSpace(v)
if v == "" {
return true
}
if !strings.HasSuffix(v, "}}") || !placeholderRe.MatchString(v) {
return false
}
return credentialResidueRe.MatchString(strings.TrimSpace(placeholderRe.ReplaceAllString(v, "")))
}
// hashBase64Strings replaces long base64 JSON string values in body with
// {{sha256:<hex>}}, editing the text in place so the rest of the body keeps
// its recorded bytes.
func hashBase64Strings(body string) string {
if len(body) <= upstreamHashMin {
return body
}
v, err := decodeJSON([]byte(body))
if err != nil {
return body
}
var long []string
collectStrings(v, &long)
for _, str := range long {
prefix, payload := "", str
if strings.HasPrefix(str, "data:") {
if i := strings.Index(str, ";base64,"); i > 0 {
prefix, payload = str[:i+len(";base64,")], str[i+len(";base64,"):]
}
}
raw, ok := decodeBase64(payload)
if !ok {
continue
}
sum := sha256.Sum256(raw)
placeholder := prefix + "{{sha256:" + hex.EncodeToString(sum[:]) + "}}"
for _, enc := range jsonStringForms(str) {
body = strings.ReplaceAll(body, enc, `"`+placeholder+`"`)
}
}
return body
}
func collectStrings(v any, out *[]string) {
switch t := v.(type) {
case map[string]any:
for _, x := range t {
collectStrings(x, out)
}
case []any:
for _, x := range t {
collectStrings(x, out)
}
case string:
if len(t) > upstreamHashMin {
*out = append(*out, t)
}
}
}
// jsonStringForms returns the quoted JSON spellings a backend may have used
// for s: Go's encoding and PHP's, which also escapes "/".
func jsonStringForms(s string) []string {
raw, _ := json.Marshal(s)
forms := []string{string(raw), `"` + s + `"`}
if esc := `"` + phpJSONEscape(s) + `"`; !slices.Contains(forms, esc) {
forms = append(forms, esc)
}
return forms
}