diff --git a/internal/config/config.go b/internal/config/config.go index a0368cb..50240b9 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -28,6 +28,9 @@ type RuleDTO struct { Latency string `json:"latency,omitempty"` // e.g. "250ms" Abort bool `json:"abort,omitempty"` // reset the connection HTTPStatus int `json:"httpStatus,omitempty"` // L7: synthesize this status (100..599) + + HTTPBody string `json:"httpBody,omitempty"` // L7: synthesized response body + HTTPHeaders map[string]string `json:"httpHeaders,omitempty"` // L7: added to the synthesized response } // Load reads and validates a rules file, returning the parsed fault rules. @@ -59,7 +62,7 @@ func (f File) ToRules() ([]fault.Rule, error) { func (d RuleDTO) toRule() (fault.Rule, error) { // Pass the probability pointer through unchanged: nil (unset) → always, // explicit 0 → never. fault.Match applies the unset default. - r := fault.Rule{Name: d.Name, Hosts: d.Hosts, Probability: d.Probability, Abort: d.Abort, HTTPStatus: d.HTTPStatus} + r := fault.Rule{Name: d.Name, Hosts: d.Hosts, Probability: d.Probability, Abort: d.Abort, HTTPStatus: d.HTTPStatus, HTTPBody: d.HTTPBody, HTTPHeaders: d.HTTPHeaders} if d.Probability != nil && (*d.Probability < 0 || *d.Probability > 1) { return r, fmt.Errorf("probability must be within [0,1], got %v", *d.Probability) diff --git a/internal/fault/fault.go b/internal/fault/fault.go index 195abba..7f71aed 100644 --- a/internal/fault/fault.go +++ b/internal/fault/fault.go @@ -47,6 +47,13 @@ type Rule struct { // cleartext HTTP (it is ignored on TLS/opaque connections). Selected by the // Host header, matched with the same semantics as Hosts. HTTPStatus int + + // HTTPBody, if set, replaces the default synthesized response body. + // HTTPHeaders are added to the synthesized response (Content-Length and + // Connection are always set by the proxy; Content-Type defaults but can be + // overridden here). Both apply only alongside HTTPStatus. + HTTPBody string + HTTPHeaders map[string]string } // hasL7 reports whether the rule carries an L7-only fault. @@ -54,10 +61,12 @@ func (r Rule) hasL7() bool { return r.HTTPStatus != 0 } // Action is the decision for a single connection. type Action struct { - Rule string - Latency time.Duration - Abort bool - HTTPStatus int + Rule string + Latency time.Duration + Abort bool + HTTPStatus int + HTTPBody string + HTTPHeaders map[string]string } // Engine holds an ordered rule set. The first matching rule wins. @@ -144,10 +153,12 @@ func (e *Engine) Match(dst netip.AddrPort, identity string) Action { return Action{Rule: r.Name} } return Action{ - Rule: r.Name, - Latency: r.Latency, - Abort: r.Abort, - HTTPStatus: r.HTTPStatus, + Rule: r.Name, + Latency: r.Latency, + Abort: r.Abort, + HTTPStatus: r.HTTPStatus, + HTTPBody: r.HTTPBody, + HTTPHeaders: r.HTTPHeaders, } } return Action{} diff --git a/internal/proxy/http.go b/internal/proxy/http.go index 0ce0a9d..5d10a9e 100644 --- a/internal/proxy/http.go +++ b/internal/proxy/http.go @@ -9,6 +9,8 @@ import ( "io" "net" "net/http" + "net/textproto" + "strconv" "strings" ) @@ -126,14 +128,35 @@ func isHTTPMethodStart(b byte) bool { // writeHTTPStatus synthesizes a minimal HTTP/1.1 response with the given status // code — an injected fault that never reaches the upstream. -func writeHTTPStatus(c net.Conn, status int) error { +// writeHTTPResponse synthesizes a cleartext HTTP response. body, if empty, +// falls back to a default one-liner. Caller headers are added as-is (their +// Content-Type overrides the default); Content-Length and Connection are always +// set by the proxy so they stay correct and the connection closes cleanly. +func writeHTTPResponse(c net.Conn, status int, headers map[string]string, body string) error { reason := http.StatusText(status) if reason == "" { reason = "Fault Injected" } - body := fmt.Sprintf("%d %s (injected by steadybit transparent-proxy)\n", status, reason) - _, err := fmt.Fprintf(c, - "HTTP/1.1 %d %s\r\nContent-Type: text/plain; charset=utf-8\r\nContent-Length: %d\r\nConnection: close\r\n\r\n%s", - status, reason, len(body), body) + if body == "" { + body = fmt.Sprintf("%d %s (injected by steadybit transparent-proxy)\n", status, reason) + } + + h := map[string]string{"Content-Type": "text/plain; charset=utf-8"} + for k, v := range headers { + h[textproto.CanonicalMIMEHeaderKey(k)] = v + } + // Proxy-owned headers: keep the framing correct regardless of caller input. + h["Content-Length"] = strconv.Itoa(len(body)) + h["Connection"] = "close" + + var sb strings.Builder + fmt.Fprintf(&sb, "HTTP/1.1 %d %s\r\n", status, reason) + for k, v := range h { + fmt.Fprintf(&sb, "%s: %s\r\n", k, v) + } + sb.WriteString("\r\n") + sb.WriteString(body) + + _, err := c.Write([]byte(sb.String())) return err } diff --git a/internal/proxy/http_test.go b/internal/proxy/http_test.go new file mode 100644 index 0000000..70077d1 --- /dev/null +++ b/internal/proxy/http_test.go @@ -0,0 +1,78 @@ +// SPDX-License-Identifier: MIT +// SPDX-FileCopyrightText: 2026 Steadybit GmbH + +package proxy + +import ( + "bufio" + "io" + "net" + "net/http" + "strconv" + "testing" +) + +// readSynthesized runs writeHTTPResponse over an in-memory pipe and parses the +// result back into an *http.Response. +func readSynthesized(t *testing.T, status int, headers map[string]string, body string) *http.Response { + t.Helper() + srv, cli := net.Pipe() + go func() { + _ = writeHTTPResponse(srv, status, headers, body) + _ = srv.Close() + }() + resp, err := http.ReadResponse(bufio.NewReader(cli), nil) + if err != nil { + t.Fatalf("ReadResponse: %v", err) + } + return resp +} + +func Test_writeHTTPResponse_defaults(t *testing.T) { + resp := readSynthesized(t, 503, nil, "") + defer resp.Body.Close() + if resp.StatusCode != 503 { + t.Fatalf("status = %d, want 503", resp.StatusCode) + } + b, _ := io.ReadAll(resp.Body) + if len(b) == 0 { + t.Fatal("expected a default body") + } + if got := resp.Header.Get("Content-Type"); got != "text/plain; charset=utf-8" { + t.Fatalf("default Content-Type = %q", got) + } + if resp.ContentLength != int64(len(b)) { + t.Fatalf("Content-Length %d != body len %d", resp.ContentLength, len(b)) + } +} + +func Test_writeHTTPResponse_customBodyAndHeaders(t *testing.T) { + body := `{"error":"nope"}` + resp := readSynthesized(t, 429, map[string]string{ + "content-type": "application/json", // lower-case, should be canonicalized + override default + "Retry-After": "30", + "X-Fault": "injected", + }, body) + defer resp.Body.Close() + + if resp.StatusCode != 429 { + t.Fatalf("status = %d, want 429", resp.StatusCode) + } + if got := resp.Header.Get("Content-Type"); got != "application/json" { + t.Fatalf("Content-Type = %q, want application/json", got) + } + if got := resp.Header.Get("Retry-After"); got != "30" { + t.Fatalf("Retry-After = %q", got) + } + if got := resp.Header.Get("X-Fault"); got != "injected" { + t.Fatalf("X-Fault = %q", got) + } + b, _ := io.ReadAll(resp.Body) + if string(b) != body { + t.Fatalf("body = %q, want %q", b, body) + } + // Content-Length is proxy-owned and must match the custom body. + if got := resp.Header.Get("Content-Length"); got != strconv.Itoa(len(body)) { + t.Fatalf("Content-Length = %q, want %d", got, len(body)) + } +} diff --git a/internal/proxy/server.go b/internal/proxy/server.go index 10a6e1f..945a81e 100644 --- a/internal/proxy/server.go +++ b/internal/proxy/server.go @@ -228,7 +228,7 @@ func (s *Server) handle(ctx context.Context, client *net.TCPConn) { // L7: synthesize an HTTP status response without contacting the upstream. // Only valid for cleartext HTTP; ignored otherwise. if proto == protoHTTP && action.HTTPStatus != 0 { - if err := writeHTTPStatus(client, action.HTTPStatus); err != nil { + if err := writeHTTPResponse(client, action.HTTPStatus, action.HTTPHeaders, action.HTTPBody); err != nil { log.Debug("failed to write injected status", slog.Any("err", err)) } s.Metrics.Proxied() diff --git a/main.go b/main.go index 7ec47a9..a12d46d 100644 --- a/main.go +++ b/main.go @@ -62,10 +62,13 @@ func main() { faultLatency = flag.Duration("fault-latency", 0, "single fault: latency added before connecting upstream") faultReset = flag.Bool("fault-reset", false, "single fault: reset (RST) matching connections") faultStatus = flag.Int("fault-http-status", 0, "single fault: injected HTTP status (L7, cleartext HTTP)") + faultBody = flag.String("fault-http-body", "", "single fault: injected HTTP response body (L7, cleartext HTTP)") faultProb = flag.Float64("fault-probability", 1, "single fault: probability [0,1] to apply the fault per connection (default 1 = always, 0 = never)") faultHosts = flag.String("fault-hosts", "", "single fault: comma-separated host selectors (SNI/Host)") faultCIDRs = flag.String("fault-cidrs", "", "single fault: comma-separated CIDR selectors") ) + var faultHeaders stringList + flag.Var(&faultHeaders, "fault-http-header", "single fault: injected HTTP response header 'Key: Value' (repeatable)") flag.Parse() logger := slog.New(slog.NewJSONHandler(os.Stderr, &slog.HandlerOptions{Level: parseLevel(*logLevel)})) @@ -76,7 +79,7 @@ func main() { logger.Error("failed to load config", slog.Any("err", err)) os.Exit(1) } - if fr, ok, ferr := buildFlagRule(*faultLatency, *faultReset, *faultStatus, *faultProb, *faultHosts, *faultCIDRs); ferr != nil { + if fr, ok, ferr := buildFlagRule(*faultLatency, *faultReset, *faultStatus, *faultBody, faultHeaders, *faultProb, *faultHosts, *faultCIDRs); ferr != nil { logger.Error("invalid fault flags", slog.Any("err", ferr)) os.Exit(2) } else if ok { @@ -185,9 +188,35 @@ func loadRules(path string) ([]fault.Rule, error) { return config.Load(path) } +// stringList is a repeatable string flag (e.g. --fault-http-header used more +// than once). +type stringList []string + +func (s *stringList) String() string { return strings.Join(*s, ", ") } +func (s *stringList) Set(v string) error { + *s = append(*s, v) + return nil +} + +// parseHeaderList turns "Key: Value" entries into a header map. +func parseHeaderList(entries []string) (map[string]string, error) { + if len(entries) == 0 { + return nil, nil + } + h := make(map[string]string, len(entries)) + for _, e := range entries { + k, v, ok := strings.Cut(e, ":") + if k = strings.TrimSpace(k); !ok || k == "" { + return nil, fmt.Errorf("invalid header %q, want 'Key: Value'", e) + } + h[k] = strings.TrimSpace(v) + } + return h, nil +} + // buildFlagRule assembles a single fault.Rule from the --fault-* flags. ok is // false when no fault flag is set. -func buildFlagRule(latency time.Duration, reset bool, status int, probability float64, hosts, cidrs string) (fault.Rule, bool, error) { +func buildFlagRule(latency time.Duration, reset bool, status int, body string, headers []string, probability float64, hosts, cidrs string) (fault.Rule, bool, error) { if latency == 0 && !reset && status == 0 { return fault.Rule{}, false, nil } @@ -197,8 +226,12 @@ func buildFlagRule(latency time.Duration, reset bool, status int, probability fl if status != 0 && (status < 100 || status > 599) { return fault.Rule{}, false, fmt.Errorf("fault-http-status must be within [100,599], got %d", status) } + httpHeaders, err := parseHeaderList(headers) + if err != nil { + return fault.Rule{}, false, fmt.Errorf("fault-http-header: %w", err) + } // The flag default is 1.0 (always); an explicit 0 means never. - r := fault.Rule{Name: "flag-rule", Latency: latency, Abort: reset, HTTPStatus: status, Probability: &probability} + r := fault.Rule{Name: "flag-rule", Latency: latency, Abort: reset, HTTPStatus: status, HTTPBody: body, HTTPHeaders: httpHeaders, Probability: &probability} for _, h := range strings.Split(hosts, ",") { if h = strings.TrimSpace(h); h != "" { r.Hosts = append(r.Hosts, h)