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" "unicode/utf8" "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 .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:}} placeholder. const upstreamHashMin = 1024 // UpstreamResponse is the recorded vendor answer the fake replays. A body // that is not valid UTF-8 (an image, for example) is written as a YAML // !!binary scalar (base64) and read back byte for byte. type UpstreamResponse struct { Status int `yaml:"status"` Headers map[string]string `yaml:"headers,omitempty"` Body string `yaml:"body,omitempty"` } // upstreamResponseYAML is UpstreamResponse as written: the body is a plain // string, a binaryScalar, or nil when empty. type upstreamResponseYAML struct { Status int `yaml:"status"` Headers map[string]string `yaml:"headers,omitempty"` Body any `yaml:"body,omitempty"` } // MarshalYAML writes a non-UTF-8 body as !!binary so the bytes survive the // round trip. func (r UpstreamResponse) MarshalYAML() (any, error) { return upstreamResponseYAML{Status: r.Status, Headers: r.Headers, Body: yamlBody(r.Body)}, nil } // binaryScalar is a byte string written as a YAML !!binary scalar. type binaryScalar string // MarshalYAML implements yaml.BytesMarshaler. func (b binaryScalar) MarshalYAML() ([]byte, error) { return []byte("!!binary " + base64.StdEncoding.EncodeToString([]byte(b))), nil } // yamlBody is the value written for a body: nil when empty, a binaryScalar // when it is not valid UTF-8, else the string. func yamlBody(body string) any { switch { case body == "": return nil case !utf8.ValidString(body): return binaryScalar(body) } return body } // 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 := r.Body if utf8.ValidString(body) { // A binary body is replayed byte for byte, never expanded. var err error if body, err = f.store.Expand(body); 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:}} (optionally behind a data: URL prefix) matches a sent // base64 string whose decoded bytes hash to . 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:}}. // 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) if utf8.ValidString(resp.Body) { // A binary body (an image) is kept byte for byte: masking a // short variable value inside it would corrupt the file. 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:}}, 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 }