Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
45 changes: 29 additions & 16 deletions internal/interception/rules.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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
}
Expand All @@ -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
}
Expand Down
17 changes: 17 additions & 0 deletions internal/interception/rules_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
151 changes: 126 additions & 25 deletions internal/metrics/metrics.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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) {
Expand Down Expand Up @@ -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)
Expand Down
44 changes: 42 additions & 2 deletions internal/metrics/metrics_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ package metrics
import (
"encoding/json"
"net/http/httptest"
"reflect"
"sync"
"testing"
)
Expand Down Expand Up @@ -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.
Expand All @@ -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)
}
}
Expand Down
Loading
Loading