diff --git a/internal/interception/rules.go b/internal/interception/rules.go index d2a46e1..4087643 100644 --- a/internal/interception/rules.go +++ b/internal/interception/rules.go @@ -45,6 +45,11 @@ type Config struct { Mark uint32 // 0 => DefaultMark HookChains []string // 0 => ["OUTPUT"] Filter Filter + // SkipFlush omits the filter-table chain that resets already-ESTABLISHED + // flows to the target ports. By default (false) warm connection pools are + // reset so they re-establish through the proxy and immediately feel the + // fault; set it when only new connections should be affected. + SkipFlush bool } func (c Config) mark() uint32 { @@ -116,21 +121,24 @@ func (c Config) AddScript() []string { } s = append(s, "COMMIT") - // filter table: reset already-ESTABLISHED flows so pools reconnect. - s = append(s, "*filter", fmt.Sprintf(":%s - [0:0]", flush)) - s = append(s, fmt.Sprintf("-A %s -m mark --mark %s -j RETURN", flush, mark)) - for _, ex := range excludes { - s = append(s, fmt.Sprintf("-A %s -d %s -j RETURN", flush, ex)) - } - for _, in := range includes { - for _, p := range c.Filter.Ports { - s = append(s, fmt.Sprintf("-A %s -p tcp -d %s --dport %d -m conntrack --ctstate ESTABLISHED -j REJECT --reject-with tcp-reset", flush, in, p)) + // filter table: reset already-ESTABLISHED flows so pools reconnect. Skipped + // when SkipFlush is set, so only new connections are affected. + if !c.SkipFlush { + s = append(s, "*filter", fmt.Sprintf(":%s - [0:0]", flush)) + s = append(s, fmt.Sprintf("-A %s -m mark --mark %s -j RETURN", flush, mark)) + for _, ex := range excludes { + s = append(s, fmt.Sprintf("-A %s -d %s -j RETURN", flush, ex)) } + for _, in := range includes { + for _, p := range c.Filter.Ports { + s = append(s, fmt.Sprintf("-A %s -p tcp -d %s --dport %d -m conntrack --ctstate ESTABLISHED -j REJECT --reject-with tcp-reset", flush, in, p)) + } + } + for _, h := range c.hooks() { + s = append(s, fmt.Sprintf("-I %s -j %s", h, flush)) + } + s = append(s, "COMMIT") } - for _, h := range c.hooks() { - s = append(s, fmt.Sprintf("-I %s -j %s", h, flush)) - } - s = append(s, "COMMIT") return s } @@ -152,10 +160,15 @@ func (c Config) DeleteCommands() [][]string { } cmds = append(cmds, []string{"-t", "nat", "-F", redir}, []string{"-t", "nat", "-X", redir}) - for _, h := range c.hooks() { - cmds = append(cmds, []string{"-t", "filter", "-D", h, "-j", flush}) + // The filter flush chain is only installed when SkipFlush is false, so only + // then does it need removing — otherwise these deletes always fail against a + // non-existent chain and add noise that can mask a genuine cleanup failure. + if !c.SkipFlush { + for _, h := range c.hooks() { + cmds = append(cmds, []string{"-t", "filter", "-D", h, "-j", flush}) + } + cmds = append(cmds, []string{"-t", "filter", "-F", flush}, []string{"-t", "filter", "-X", flush}) } - cmds = append(cmds, []string{"-t", "filter", "-F", flush}, []string{"-t", "filter", "-X", flush}) return cmds } diff --git a/internal/interception/rules_test.go b/internal/interception/rules_test.go index 119e250..6e00d06 100644 --- a/internal/interception/rules_test.go +++ b/internal/interception/rules_test.go @@ -59,6 +59,23 @@ func TestAddScript_Structure(t *testing.T) { } } +func TestAddScript_SkipFlushOmitsFilterChain(t *testing.T) { + c := baseConfig(t) + c.SkipFlush = true + script := joined(c.AddScript()) + + // The nat REDIRECT must still be installed... + if !strings.Contains(script, "-A SB_TP_REDIR_execid123456 -p tcp -d 0.0.0.0/0 --dport 443 -j REDIRECT --to-ports 3128") { + t.Fatalf("REDIRECT missing with SkipFlush:\n%s", script) + } + // ...but the filter flush chain must be entirely absent. + for _, forbidden := range []string{"*filter", "SB_TP_FLUSH_execid123456", "ESTABLISHED"} { + if strings.Contains(script, forbidden) { + t.Errorf("SkipFlush script should not contain %q:\n%s", forbidden, script) + } + } +} + func TestAddScript_MarkExemptionFirst(t *testing.T) { c := baseConfig(t) script := c.AddScript() diff --git a/internal/metrics/metrics.go b/internal/metrics/metrics.go index 7a6255f..0a074d3 100644 --- a/internal/metrics/metrics.go +++ b/internal/metrics/metrics.go @@ -11,34 +11,62 @@ package metrics import ( "encoding/json" "net/http" + "sort" + "sync" "sync/atomic" ) // Metrics holds concurrency-safe counters shared across connection goroutines. type Metrics struct { - ConnectionsMatched atomic.Int64 // original destination resolved (candidate for fault) - ConnectionsActive atomic.Int64 // currently proxied - ConnectionsProxied atomic.Int64 // completed pass-through/forward - ConnectionsAborted atomic.Int64 // reset by an abort rule - ConnectionsDropped atomic.Int64 // loop guard / peek failure / self-refusal - UpstreamErrors atomic.Int64 // dial failures - BytesToUpstream atomic.Int64 - BytesToClient atomic.Int64 + ConnectionsMatched atomic.Int64 // original destination resolved (candidate for fault) + ConnectionsActive atomic.Int64 // currently proxied + ConnectionsProxied atomic.Int64 // completed pass-through/forward + ConnectionsAborted atomic.Int64 // reset by an abort rule + ConnectionsDropped atomic.Int64 // loop guard / peek failure / self-refusal + ConnectionsFaulted atomic.Int64 // connections a fault was actually applied to (once each) + LatencyApplied atomic.Int64 // connections a latency fault delayed + HTTPResponsesInjected atomic.Int64 // connections given a synthesized HTTP response + UpstreamErrors atomic.Int64 // dial failures + BytesToUpstream atomic.Int64 + BytesToClient atomic.Int64 + + // perHost breaks the matched/faulted counts down by the dependency hostname + // (SNI or HTTP Host) a connection carried. It is the "which dependency, how + // often" view the platform renders. Guarded by mu because hostnames are + // dynamic keys. Hostnames here are the attacker-supplied targets, so exposing + // them in the operator's own statistics is safe (unlike the process logs). + mu sync.Mutex + perHost map[string]*hostCounters +} + +type hostCounters struct { + matched int64 + faulted int64 } // New returns a ready-to-use Metrics. -func New() *Metrics { return &Metrics{} } +func New() *Metrics { return &Metrics{perHost: map[string]*hostCounters{}} } + +// HostStat is the per-hostname view in a Snapshot. +type HostStat struct { + Matched int64 `json:"matched"` + Faulted int64 `json:"faulted"` +} // Snapshot is a point-in-time, JSON-serialisable view. type Snapshot struct { - ConnectionsMatched int64 `json:"connections_matched"` - ConnectionsActive int64 `json:"connections_active"` - ConnectionsProxied int64 `json:"connections_proxied"` - ConnectionsAborted int64 `json:"connections_aborted"` - ConnectionsDropped int64 `json:"connections_dropped"` - UpstreamErrors int64 `json:"upstream_errors"` - BytesToUpstream int64 `json:"bytes_to_upstream"` - BytesToClient int64 `json:"bytes_to_client"` + ConnectionsMatched int64 `json:"connections_matched"` + ConnectionsActive int64 `json:"connections_active"` + ConnectionsProxied int64 `json:"connections_proxied"` + ConnectionsAborted int64 `json:"connections_aborted"` + ConnectionsDropped int64 `json:"connections_dropped"` + ConnectionsFaulted int64 `json:"connections_faulted"` + LatencyApplied int64 `json:"latency_applied"` + HTTPResponsesInjected int64 `json:"http_responses_injected"` + UpstreamErrors int64 `json:"upstream_errors"` + BytesToUpstream int64 `json:"bytes_to_upstream"` + BytesToClient int64 `json:"bytes_to_client"` + PerHost map[string]HostStat `json:"per_host,omitempty"` } // Snapshot reads all counters. It is not atomic across counters (values may @@ -47,18 +75,42 @@ func (m *Metrics) Snapshot() Snapshot { if m == nil { return Snapshot{} } + m.mu.Lock() + var perHost map[string]HostStat + if len(m.perHost) > 0 { + perHost = make(map[string]HostStat, len(m.perHost)) + for h, c := range m.perHost { + perHost[h] = HostStat{Matched: c.matched, Faulted: c.faulted} + } + } + m.mu.Unlock() return Snapshot{ - ConnectionsMatched: m.ConnectionsMatched.Load(), - ConnectionsActive: m.ConnectionsActive.Load(), - ConnectionsProxied: m.ConnectionsProxied.Load(), - ConnectionsAborted: m.ConnectionsAborted.Load(), - ConnectionsDropped: m.ConnectionsDropped.Load(), - UpstreamErrors: m.UpstreamErrors.Load(), - BytesToUpstream: m.BytesToUpstream.Load(), - BytesToClient: m.BytesToClient.Load(), + ConnectionsMatched: m.ConnectionsMatched.Load(), + ConnectionsActive: m.ConnectionsActive.Load(), + ConnectionsProxied: m.ConnectionsProxied.Load(), + ConnectionsAborted: m.ConnectionsAborted.Load(), + ConnectionsDropped: m.ConnectionsDropped.Load(), + ConnectionsFaulted: m.ConnectionsFaulted.Load(), + LatencyApplied: m.LatencyApplied.Load(), + HTTPResponsesInjected: m.HTTPResponsesInjected.Load(), + UpstreamErrors: m.UpstreamErrors.Load(), + BytesToUpstream: m.BytesToUpstream.Load(), + BytesToClient: m.BytesToClient.Load(), + PerHost: perHost, } } +// SortedHosts returns the per-host keys in a stable order, for deterministic +// rendering. +func (s Snapshot) SortedHosts() []string { + hosts := make([]string, 0, len(s.PerHost)) + for h := range s.PerHost { + hosts = append(hosts, h) + } + sort.Strings(hosts) + return hosts +} + // Handler serves the current snapshot as JSON at any path. func (m *Metrics) Handler() http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { @@ -101,6 +153,55 @@ func (m *Metrics) Proxied() { } } +// Faulted records that a fault was actually applied to a connection. Call it at +// most once per connection (a connection can carry both latency and an injected +// response, but is a single faulted connection). +func (m *Metrics) Faulted() { + if m != nil { + m.ConnectionsFaulted.Add(1) + } +} + +// LatencyInjected records a connection a latency fault delayed. +func (m *Metrics) LatencyInjected() { + if m != nil { + m.LatencyApplied.Add(1) + } +} + +// HTTPInjected records a connection given a synthesized HTTP response (instead +// of being forwarded upstream). +func (m *Metrics) HTTPInjected() { + if m != nil { + m.HTTPResponsesInjected.Add(1) + } +} + +// MatchedHost records that a connection carrying the given dependency hostname +// matched a rule. FaultedHost records that a fault was actually applied to it +// (i.e. it passed the probability roll). Empty hosts are ignored. +func (m *Metrics) MatchedHost(host string) { m.addHost(host, true, false) } +func (m *Metrics) FaultedHost(host string) { m.addHost(host, false, true) } + +func (m *Metrics) addHost(host string, matched, faulted bool) { + if m == nil || host == "" { + return + } + m.mu.Lock() + defer m.mu.Unlock() + c := m.perHost[host] + if c == nil { + c = &hostCounters{} + m.perHost[host] = c + } + if matched { + c.matched++ + } + if faulted { + c.faulted++ + } +} + func (m *Metrics) UpstreamError() { if m != nil { m.UpstreamErrors.Add(1) diff --git a/internal/metrics/metrics_test.go b/internal/metrics/metrics_test.go index 8ac1b44..d64c8fb 100644 --- a/internal/metrics/metrics_test.go +++ b/internal/metrics/metrics_test.go @@ -6,6 +6,7 @@ package metrics import ( "encoding/json" "net/http/httptest" + "reflect" "sync" "testing" ) @@ -43,11 +44,50 @@ func TestMetrics_CountersAndActiveGauge(t *testing.T) { BytesToUpstream: 100, BytesToClient: 250, } - if s != want { + if !reflect.DeepEqual(s, want) { t.Fatalf("snapshot = %+v, want %+v", s, want) } } +func TestMetrics_PerHostAndFaultCounters(t *testing.T) { + m := New() + + // two connections to api.example.com, one faulted; one to cdn.example.com, + // faulted with an injected HTTP response and a latency. + m.MatchedHost("api.example.com") + m.MatchedHost("api.example.com") + m.FaultedHost("api.example.com") + m.Faulted() + m.MatchedHost("cdn.example.com") + m.FaultedHost("cdn.example.com") + m.Faulted() + // one connection carried both latency and an injected response, but is a + // single faulted connection. + m.LatencyInjected() + m.HTTPInjected() + m.MatchedHost("") // ignored + + s := m.Snapshot() + if s.LatencyApplied != 1 || s.HTTPResponsesInjected != 1 { + t.Fatalf("fault counters = %+v", s) + } + if s.ConnectionsFaulted != 2 { + t.Fatalf("ConnectionsFaulted = %d, want 2 (once per faulted connection)", s.ConnectionsFaulted) + } + if got := s.PerHost["api.example.com"]; got.Matched != 2 || got.Faulted != 1 { + t.Fatalf("api per-host = %+v", got) + } + if got := s.PerHost["cdn.example.com"]; got.Matched != 1 || got.Faulted != 1 { + t.Fatalf("cdn per-host = %+v", got) + } + if _, ok := s.PerHost[""]; ok { + t.Fatalf("empty host should not be recorded") + } + if hosts := s.SortedHosts(); len(hosts) != 2 || hosts[0] != "api.example.com" { + t.Fatalf("sorted hosts = %v", hosts) + } +} + func TestMetrics_NilSafe(t *testing.T) { var m *Metrics // nil // None of these must panic. @@ -57,7 +97,7 @@ func TestMetrics_NilSafe(t *testing.T) { m.Proxied() m.UpstreamError() m.AddBytes(1, 2) - if s := m.Snapshot(); (s != Snapshot{}) { + if s := m.Snapshot(); !reflect.DeepEqual(s, Snapshot{}) { t.Fatalf("nil metrics should snapshot to zero, got %+v", s) } } diff --git a/internal/proxy/server.go b/internal/proxy/server.go index 945a81e..77de879 100644 --- a/internal/proxy/server.go +++ b/internal/proxy/server.go @@ -205,11 +205,32 @@ func (s *Server) handle(ctx context.Context, client *net.TCPConn) { action := s.Faults.Match(dst, identity) log = log.With( slog.String("dst", dst.String()), - slog.String("identity", identity), slog.String("rule", action.Rule), ) + // identity (the TLS SNI or HTTP Host) is potentially sensitive, so it is + // only ever logged at debug — never at the info level that ships by default + // — while still feeding the per-host statistics the operator sees. + if identity != "" { + log.Debug("matched dependency", slog.String("identity", identity)) + } + if action.Rule != "" { + s.Metrics.MatchedHost(identity) + } + // faulted is recorded at the point a fault actually applies — not + // speculatively from the Action — so per-host and global "faulted" counts + // stay in step and never over-report (e.g. an HTTP rule on a TLS connection, + // which is forwarded untouched). It fires at most once per connection. + faultRecorded := false + markFaulted := func() { + if !faultRecorded { + faultRecorded = true + s.Metrics.Faulted() + s.Metrics.FaultedHost(identity) + } + } if action.Abort { + markFaulted() s.Metrics.Aborted() log.Info("aborting connection (reset)") reset(client) @@ -217,6 +238,8 @@ func (s *Server) handle(ctx context.Context, client *net.TCPConn) { } if action.Latency > 0 { + markFaulted() + s.Metrics.LatencyInjected() select { case <-time.After(action.Latency): case <-ctx.Done(): @@ -228,10 +251,11 @@ 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 { + markFaulted() 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() + s.Metrics.HTTPInjected() log.Info("injected http status", slog.Int("status", action.HTTPStatus)) return } diff --git a/main.go b/main.go index a12d46d..a2b1dbe 100644 --- a/main.go +++ b/main.go @@ -15,6 +15,7 @@ package main import ( "context" + "encoding/json" "errors" "flag" "fmt" @@ -40,13 +41,15 @@ import ( func main() { var ( - listen = flag.String("listen", "0.0.0.0:3128", "address to accept redirected connections on") - rulesPath = flag.String("config", "", "path to a JSON fault-rules file (optional; empty = pure pass-through)") - logLevel = flag.String("log-level", "info", "log level: debug, info, warn, error") - dialTimeout = flag.Duration("dial-timeout", 10*time.Second, "upstream connection timeout") - mark = flag.Uint("mark", uint(interception.DefaultMark), "SO_MARK stamped on upstream sockets for interception loop-protection (0 disables)") - metricsAddr = flag.String("metrics-addr", "", "address to serve JSON metrics on (empty disables)") - maxDuration = flag.Duration("max-duration", 0, "deadman: self-terminate and tear down after this long (0 disables)") + listen = flag.String("listen", "0.0.0.0:3128", "address to accept redirected connections on") + rulesPath = flag.String("config", "", "path to a JSON fault-rules file (optional; empty = pure pass-through)") + logLevel = flag.String("log-level", "info", "log level: debug, info, warn, error") + dialTimeout = flag.Duration("dial-timeout", 10*time.Second, "upstream connection timeout") + mark = flag.Uint("mark", uint(interception.DefaultMark), "SO_MARK stamped on upstream sockets for interception loop-protection (0 disables)") + metricsAddr = flag.String("metrics-addr", "", "address to serve JSON metrics on (empty disables)") + metricsStdout = flag.Duration("metrics-stdout-interval", 0, "if >0, print a JSON metrics snapshot to stdout at this interval (and once on exit)") + maxDuration = flag.Duration("max-duration", 0, "deadman: self-terminate and tear down after this long (0 disables)") + noFlush = flag.Bool("no-flush", false, "do not reset already-ESTABLISHED connections on start (only new connections are affected)") prePorts = flag.String("preflight-ports", "", "comma-separated target ports to refuse-on-conflict against an existing mesh (defaults to --intercept-ports)") @@ -106,11 +109,17 @@ func main() { if *metricsAddr != "" { serveMetrics(ctx, logger, *metricsAddr, m) } + // Periodic metrics on stdout — the cross-namespace-friendly channel the + // orchestrating extension scrapes (an HTTP endpoint in a container's netns is + // unreachable from the extension; stdout is always captured). + if *metricsStdout > 0 { + streamMetricsToStdout(ctx, *metricsStdout, m) + } // Validate the interception filter and whether self-managed mode is wanted. // The port is filled in after we bind (below); 0 is fine here because this // instance is only used for --revert, where the port is irrelevant. - interceptor, wantIntercept, err := buildInterceptor(*interceptCIDRs, *interceptPorts, *excludeCIDRs, *execID, uint32(*mark), 0) + interceptor, wantIntercept, err := buildInterceptor(*interceptCIDRs, *interceptPorts, *excludeCIDRs, *execID, uint32(*mark), 0, *noFlush) if err != nil { logger.Error("invalid interception configuration", slog.Any("err", err)) os.Exit(2) @@ -155,7 +164,7 @@ func main() { os.Exit(1) } port := uint16(ln.Addr().(*net.TCPAddr).Port) - interceptor, _, err = buildInterceptor(*interceptCIDRs, *interceptPorts, *excludeCIDRs, *execID, uint32(*mark), port) + interceptor, _, err = buildInterceptor(*interceptCIDRs, *interceptPorts, *excludeCIDRs, *execID, uint32(*mark), port, *noFlush) if err != nil { logger.Error("invalid interception configuration", slog.Any("err", err)) os.Exit(2) @@ -192,7 +201,7 @@ func loadRules(path string) ([]fault.Rule, error) { // than once). type stringList []string -func (s *stringList) String() string { return strings.Join(*s, ", ") } +func (s *stringList) String() string { return strings.Join(*s, ", ") } func (s *stringList) Set(v string) error { *s = append(*s, v) return nil @@ -259,7 +268,7 @@ func (a interceptorAdapter) Revert(ctx context.Context) error { return a.cfg.Rev // targets the exact port the proxy bound — there is no pre-allocated port that // another process could steal between allocation and bind. For --revert the // port is irrelevant (chain names derive from the exec-id) and may be 0. -func buildInterceptor(cidrs, ports, excludes, execID string, mark uint32, proxyPort uint16) (supervisor.Interceptor, bool, error) { +func buildInterceptor(cidrs, ports, excludes, execID string, mark uint32, proxyPort uint16, noFlush bool) (supervisor.Interceptor, bool, error) { if cidrs == "" && ports == "" { return nil, false, nil } @@ -284,6 +293,7 @@ func buildInterceptor(cidrs, ports, excludes, execID string, mark uint32, proxyP ExecutionID: execID, ProxyPort: proxyPort, Mark: mark, + SkipFlush: noFlush, Filter: interception.Filter{ Include: include, Exclude: exclude, @@ -325,6 +335,27 @@ func serveMetrics(ctx context.Context, logger *slog.Logger, addr string, m *metr logger.Info("serving metrics", slog.String("addr", addr)) } +// streamMetricsToStdout prints a JSON metrics snapshot (one compact line) to +// stdout every interval, plus a final line when the context is cancelled. The +// proxy's own structured logs go to stderr, so stdout carries only these +// snapshots and the orchestrating extension can scrape it line by line. +func streamMetricsToStdout(ctx context.Context, interval time.Duration, m *metrics.Metrics) { + enc := json.NewEncoder(os.Stdout) + go func() { + t := time.NewTicker(interval) + defer t.Stop() + for { + select { + case <-ctx.Done(): + _ = enc.Encode(m.Snapshot()) // final snapshot + return + case <-t.C: + _ = enc.Encode(m.Snapshot()) + } + } + }() +} + func parseCIDRs(s string) ([]netip.Prefix, error) { var out []netip.Prefix for _, c := range strings.Split(s, ",") {