From 76056dd2e56f04d1c1984160f54098e403fbb718 Mon Sep 17 00:00:00 2001 From: rpb-ant Date: Mon, 27 Jul 2026 05:13:25 -0400 Subject: [PATCH] plugin/shed: add UDP overload protection plugin (#8312) * plugin/shed: add UDP overload protection plugin UDP responses written back through one listener socket serialize on the Go runtime's internal fdMutex, which allows at most 2^20-1 concurrent operations per file descriptor and panics the process when exceeded. CoreDNS serves UDP with one goroutine per query, all writing through the shared packet connection, so a sustained overload parks every excess in-flight query in that wait queue until the process dies with "too many concurrent operations on a single file or socket". Observed in production: ~2.8M goroutines and 60GiB RSS before the panic. The shed plugin makes the panic structurally unreachable. It installs, via Config.UDPDecorateWriterFunc, a per-socket bounded evict-oldest stack drained newest-first by a single writer goroutine, so the fd never sees more than one writer and residual capacity under overload always goes to the freshest response. While a socket's stack is full, arriving queries are dropped before any plugin runs. Drops are silent (the client's resolver retries elsewhere) and counted in coredns_shed_dropped_total{server, reason}. plugin/shed/fdmutex_test.go demonstrates the failure and the fix with one shared flood harness. Two subprocess tests reproduce the exact runtime panic without the plugin's write discipline - one deterministic (a held write plus >2^20 queued writers), one with nothing held or mocked; both exercise the Go runtime rather than the plugin, so they are gated behind SHED_FLOOD_TEST=1. The counterfactual - the same load through the plugin's stack, completing with every response accounted for as written or dropped - runs in every test invocation, including -race, at 50k responders, and at the full 1.5M with SHED_FLOOD_TEST=1: SHED_FLOOD_TEST=1 go test ./plugin/shed/ Signed-off-by: Ryan Brewster * test: add shed e2e test Query a shed-enabled server over UDP (the plugin's deferred single-writer path) and TCP (which shed passes through), and check that coredns_shed_dropped_total is exported with its reason label. No-Verification-Needed: test-only change Signed-off-by: Ryan Brewster --------- Signed-off-by: Ryan Brewster --- core/dnsserver/zdirectives.go | 1 + core/plugin/zplugin.go | 1 + plugin.cfg | 1 + plugin/shed/README.md | 79 ++++++++++ plugin/shed/fdmutex_test.go | 273 ++++++++++++++++++++++++++++++++++ plugin/shed/metrics.go | 15 ++ plugin/shed/setup.go | 68 +++++++++ plugin/shed/setup_test.go | 53 +++++++ plugin/shed/shed.go | 116 +++++++++++++++ plugin/shed/shed_test.go | 172 +++++++++++++++++++++ plugin/shed/stack.go | 176 ++++++++++++++++++++++ plugin/shed/stack_test.go | 164 ++++++++++++++++++++ test/shed_test.go | 47 ++++++ 13 files changed, 1166 insertions(+) create mode 100644 plugin/shed/README.md create mode 100644 plugin/shed/fdmutex_test.go create mode 100644 plugin/shed/metrics.go create mode 100644 plugin/shed/setup.go create mode 100644 plugin/shed/setup_test.go create mode 100644 plugin/shed/shed.go create mode 100644 plugin/shed/shed_test.go create mode 100644 plugin/shed/stack.go create mode 100644 plugin/shed/stack_test.go create mode 100644 test/shed_test.go diff --git a/core/dnsserver/zdirectives.go b/core/dnsserver/zdirectives.go index ec8d654d7..f6549ff75 100644 --- a/core/dnsserver/zdirectives.go +++ b/core/dnsserver/zdirectives.go @@ -31,6 +31,7 @@ var Directives = []string{ "ready", "health", "pprof", + "shed", "prometheus", "errors", "log", diff --git a/core/plugin/zplugin.go b/core/plugin/zplugin.go index 476d8c1dc..1efb33f62 100644 --- a/core/plugin/zplugin.go +++ b/core/plugin/zplugin.go @@ -54,6 +54,7 @@ import ( _ "github.com/coredns/coredns/plugin/root" _ "github.com/coredns/coredns/plugin/route53" _ "github.com/coredns/coredns/plugin/secondary" + _ "github.com/coredns/coredns/plugin/shed" _ "github.com/coredns/coredns/plugin/sign" _ "github.com/coredns/coredns/plugin/template" _ "github.com/coredns/coredns/plugin/timeouts" diff --git a/plugin.cfg b/plugin.cfg index 055b08bdb..b971b5e08 100644 --- a/plugin.cfg +++ b/plugin.cfg @@ -40,6 +40,7 @@ trace:trace ready:ready health:health pprof:pprof +shed:shed prometheus:metrics errors:errors log:log diff --git a/plugin/shed/README.md b/plugin/shed/README.md new file mode 100644 index 000000000..02fa1e4b6 --- /dev/null +++ b/plugin/shed/README.md @@ -0,0 +1,79 @@ +# shed + +## Name + +*shed* - serializes UDP response writes per listener socket and sheds load when the socket cannot keep up. + +## Description + +UDP responses written back through one listener socket serialize on the Go runtime's internal +fdMutex, which allows at most 2^20-1 concurrent operations (holders plus waiters) per file +descriptor and terminates the process with + +~~~ txt +panic: too many concurrent operations on a single file or socket (max 1048575) +~~~ + +when that is exceeded. CoreDNS serves UDP with one goroutine per query, all writing back through +the shared packet connection, so when queries arrive faster than the socket's serialized writes +drain, every excess in-flight query parks its goroutine in that wait queue and nothing bounds the +pile. Observed in production: ~2.8M goroutines and 60GiB RSS before the panic. + +The *shed* plugin makes that panic structurally unreachable, per UDP listener socket: + +* **Single writer** - responses are not written by the handler goroutine. The packed response is + pushed onto a bounded per-socket stack (fixed depth 1024) and one writer goroutine per socket + performs the wire writes, so the file descriptor never sees more than one writer. The stack + evicts the oldest entry when full and the writer pops the newest first, so under overload the + socket's residual capacity always goes to the freshest response. The depth is a fixed burst + budget (roughly 12-16ms of a typical socket's drain rate), not a tunable. +* **Coupled shedding** - while a socket's stack is full, arriving queries on that socket are + dropped before any plugin runs; work admitted then would only produce a response destined for + eviction. There is no configuration: the stack's fullness is the signal. + +Drops are silent - no response is written, so the client's resolver retries against another +server, the standard load-shedding contract for UDP DNS. Every drop is counted. + +The plugin only acts on UDP; TCP queries pass through untouched. It can only be used in plain DNS +server blocks (not *tls*, *grpc*, *https* or *quic*), which is enforced at startup. It should be +listed before (above) the *prometheus* plugin in the plugin chain, so that shed drops are never +counted as handled requests by the *prometheus* plugin - which is where this plugin sits by +default. + +When several server blocks share a listener, any block with *shed* installs the write discipline +for every write on that socket, while the pre-chain shedding only runs in blocks that carry the +directive - keep it uniform across blocks sharing a listener. The discipline covers every response +written through `WriteMsg`, which is how every plugin responds; a plugin writing raw bytes with +`ResponseWriter.Write` would bypass it. + +## Syntax + +~~~ txt +shed +~~~ + +The plugin takes no arguments. + +## Metrics + +If monitoring is enabled (via the *prometheus* plugin) then the following metric is exported: + +* `coredns_shed_dropped_total{server, reason}` - counter of dropped queries and responses. The + `reason` label is `query` for queries dropped before the plugin chain because the socket's + stack was full, and `response` for responses dropped at the write boundary (evicted by a newer + response, failed to reach the wire, or arriving during shutdown). + +## Examples + +Protect the UDP listener while forwarding: + +~~~ corefile +. { + shed + forward . 8.8.8.8 +} +~~~ + +## See Also + +The fdMutex limit is enforced in `GOROOT/src/internal/poll/fd_mutex.go`. diff --git a/plugin/shed/fdmutex_test.go b/plugin/shed/fdmutex_test.go new file mode 100644 index 000000000..cc5af7e07 --- /dev/null +++ b/plugin/shed/fdmutex_test.go @@ -0,0 +1,273 @@ +package shed + +// Evidence tests for the fdMutex overflow panic described in README.md. +// One flood harness, two write disciplines: +// +// - TestFdMutexPanicOneSlowWrite (SHED_FLOOD_TEST=1): one held write plus +// >2^20 queued raw writers deterministically panic a subprocess. +// - TestFdMutexPanicUDPFlood (SHED_FLOOD_TEST=1): the same panic with +// nothing held — raw writers simply outpace the serialized drain. +// - TestSingleWriterNoPanicSameLoad: the same responders through the +// plugin's stack and single writer complete with every response written +// or counted dropped. Runs at 50k responders by default (including +// -race in CI); at the full 1.5M under SHED_FLOOD_TEST=1. +// +// The panic tests re-exec the test binary (the panic is a process death), +// cost ~1.5M goroutines / a few GiB / seconds, and assert on the runtime's +// message — env-gated so no automated or casual run pays that, or breaks if +// a future Go release rewords the panic. + +import ( + "fmt" + "net" + "os" + "os/exec" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/prometheus/client_golang/prometheus/testutil" +) + +// overflowMsg must match GOROOT/src/internal/poll/fd_mutex.go. +const overflowMsg = "too many concurrent operations on a single file or socket (max 1048575)" + +const ( + childEnv = "COREDNS_SHED_FDMUTEX_CHILD" // "flood", "held", absent = normal run + floodEnv = "SHED_FLOOD_TEST" // set to run the panic tests and the full-size survival test + + // fdMutex fields are 20-bit: the 1,048,576th concurrent op panics. + fdMutexLimit = 1 << 20 + + // floodWriters is comfortably above the limit, so the flood mode still + // crosses it after subtracting whatever the drain completes while + // spawning. ciWriters exercises the same code paths at a size every + // test run can afford. + floodWriters = 1_500_000 + ciWriters = 50_000 + nSpawners = 16 + + // Flood mode uses near-max UDP payloads so each serialized sendto is + // expensive — a stand-in for a response datapath slower than the + // arrival rate. + floodPayload = 63 * 1024 + + childTimeout = 120 * time.Second +) + +func TestMain(m *testing.M) { + switch os.Getenv(childEnv) { + case "flood": + childFlood(false) + case "held": + childFlood(true) + default: + os.Exit(m.Run()) + } + // The parent asserts on the exit status; this line is log-only. + fmt.Println("CHILD-SURVIVED-WITHOUT-PANIC") + os.Exit(0) +} + +// spawnResponders spawns n goroutines, each calling respond once — a raw +// socket write in the panic modes, a stack push in the survival mode. +func spawnResponders(n int, respond func()) (started, completed *atomic.Int64) { + started, completed = new(atomic.Int64), new(atomic.Int64) + var spawn sync.WaitGroup + for range nSpawners { + spawn.Go(func() { + for range n / nSpawners { + started.Add(1) + go func() { + respond() + completed.Add(1) + }() + } + }) + } + spawn.Wait() + return started, completed +} + +// childFlood is the crash payload: pile >2^20 concurrent raw writes onto one +// UDP socket. With held=true, one in-progress write is first parked via +// SyscallConn so the pile-up is deterministic; with held=false the writers +// race a genuine serialized drain. +func childFlood(held bool) { + conn, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + fmt.Println("child: listen:", err) + return + } + sink, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + fmt.Println("child: sink listen:", err) + return + } + dst := sink.LocalAddr().(*net.UDPAddr) + + payload := make([]byte, floodPayload) + if held { + payload = payload[:64] // writes only queue as waiters; size is irrelevant + // Park one write in progress: the callback holds the fd's write + // lock exactly as a write blocked in the kernel would. Everything + // arriving behind it becomes an fdMutex waiter. + rc, err := conn.SyscallConn() + if err != nil { + fmt.Println("child: syscallconn:", err) + return + } + holding := make(chan struct{}) + go func() { + rc.Write(func(uintptr) bool { + close(holding) + select {} // hold the write lock for the life of the process + }) + }() + <-holding + fmt.Println("child: one slow write in progress (fd write lock held)") + } + + fmt.Printf("child: spawning %d concurrent UDP writers on one socket (limit %d)\n", + floodWriters, fdMutexLimit-1) + started, completed := spawnResponders(floodWriters, func() { + conn.WriteToUDP(payload, dst) //nolint:errcheck // the pile, not the result, is the point + }) + + // If the panic is going to happen it already has (it fires inside a + // writer's WriteToUDP). Give the drain a moment, then report survival. + deadline := time.Now().Add(childTimeout) + for completed.Load() < started.Load() && time.Now().Before(deadline) { + time.Sleep(100 * time.Millisecond) + } +} + +// runCrashChild re-execs this test binary in the given child mode and +// returns its combined output. The child is expected to die. +func runCrashChild(t *testing.T, mode string) string { + t.Helper() + cmd := exec.Command(os.Args[0], "-test.run=^$") + cmd.Env = append(os.Environ(), childEnv+"="+mode, "GOTRACEBACK=single") + start := time.Now() + out, err := cmd.CombinedOutput() + t.Logf("child (%s) ran %v, err=%v", mode, time.Since(start).Round(time.Millisecond), err) + s := string(out) + // Panic output ends with a goroutine stack; keep the log readable. + if i := strings.Index(s, "goroutine "); i > 0 { + t.Logf("child output:\n%s[stack trace elided]", s[:i]) + } else { + t.Logf("child output:\n%s", s) + } + if err == nil { + t.Fatal("child process survived — expected fdMutex overflow panic") + } + return s +} + +func skipUnlessFloodTest(t *testing.T) { + t.Helper() + if os.Getenv(floodEnv) == "" { + t.Skipf("panic reproduction (~%d goroutines, a few GiB); set %s=1 to run", floodWriters, floodEnv) + } +} + +// TestFdMutexPanicUDPFlood: >2^20 genuinely concurrent writes on one UDP +// socket kill the process. Nothing is held or mocked — the writers simply +// arrive faster than the fd's serialized writes drain, which is the +// production storm condition. +func TestFdMutexPanicUDPFlood(t *testing.T) { + skipUnlessFloodTest(t) + out := runCrashChild(t, "flood") + if !strings.Contains(out, overflowMsg) { + t.Fatalf("child died without the fdMutex overflow panic; want %q", overflowMsg) + } +} + +// TestFdMutexPanicOneSlowWrite: deterministic variant — a single slow +// in-progress write plus >2^20 queued writers overflow the fdMutex waiter +// counter. No timing or throughput assumptions. +func TestFdMutexPanicOneSlowWrite(t *testing.T) { + skipUnlessFloodTest(t) + out := runCrashChild(t, "held") + if !strings.Contains(out, overflowMsg) { + t.Fatalf("child died without the fdMutex overflow panic; want %q", overflowMsg) + } +} + +// countingUDPWriter is the dns.Writer handed to the plugin's stack: the raw +// wire write, counted on success (a failed write is counted by the plugin +// as a drop). +type countingUDPWriter struct { + conn *net.UDPConn + dst *net.UDPAddr + written atomic.Int64 +} + +func (w *countingUDPWriter) Write(p []byte) (int, error) { + n, err := w.conn.WriteToUDP(p, w.dst) + if err == nil { + w.written.Add(1) + } + return n, err +} + +// TestSingleWriterNoPanicSameLoad drives the flood harness through the +// plugin's actual respStack and writer goroutine. Only that one goroutine +// ever touches the fd, so the fdMutex overflow is structurally unreachable, +// and every response is accounted for as written or dropped. +func TestSingleWriterNoPanicSameLoad(t *testing.T) { + n := ciWriters + if os.Getenv(floodEnv) != "" { + n = floodWriters + } + + conn, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatal(err) + } + defer conn.Close() + sink, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatal(err) + } + defer sink.Close() + + dropped := droppedTotal.WithLabelValues(t.Name(), "response") + droppedBefore := testutil.ToFloat64(dropped) // the child accumulates across -count>1 runs + rs := newRespStack(stackDepth, dropped) + writerDone := make(chan struct{}) + go func() { + defer close(writerDone) + rs.writerLoop() + }() + + w := &countingUDPWriter{conn: conn, dst: sink.LocalAddr().(*net.UDPAddr)} + payload := make([]byte, 64) + + start := time.Now() + started, completed := spawnResponders(n, func() { + // The responder's entire write path: what the decorator installs. + (&stackWriter{stack: rs, inner: w}).Write(payload) //nolint:errcheck // always reports success + }) + deadline := time.Now().Add(childTimeout) + for completed.Load() < started.Load() { + if time.Now().After(deadline) { + t.Fatalf("only %d/%d responders completed", completed.Load(), started.Load()) + } + time.Sleep(10 * time.Millisecond) + } + elapsed := time.Since(start) + + rs.close() + <-writerDone + + written := w.written.Load() + droppedN := int64(testutil.ToFloat64(dropped) - droppedBefore) + if spawned := started.Load(); written+droppedN != spawned { + t.Fatalf("accounting: written=%d + dropped=%d != %d responders", written, droppedN, spawned) + } + t.Logf("%d concurrent responders completed in %v with ONE fd writer: %d responses written, %d evicted (counted drops), no panic", + started.Load(), elapsed.Round(time.Millisecond), written, droppedN) +} diff --git a/plugin/shed/metrics.go b/plugin/shed/metrics.go new file mode 100644 index 000000000..d69a7677e --- /dev/null +++ b/plugin/shed/metrics.go @@ -0,0 +1,15 @@ +package shed + +import ( + "github.com/coredns/coredns/plugin" + + "github.com/prometheus/client_golang/prometheus" + "github.com/prometheus/client_golang/prometheus/promauto" +) + +var droppedTotal = promauto.NewCounterVec(prometheus.CounterOpts{ + Namespace: plugin.Namespace, + Subsystem: "shed", + Name: "dropped_total", + Help: "Counter of queries and responses dropped, per server, by reason.", +}, []string{"server", "reason"}) diff --git a/plugin/shed/setup.go b/plugin/shed/setup.go new file mode 100644 index 000000000..4b2713007 --- /dev/null +++ b/plugin/shed/setup.go @@ -0,0 +1,68 @@ +package shed + +import ( + "github.com/coredns/caddy" + "github.com/coredns/coredns/core/dnsserver" + "github.com/coredns/coredns/plugin" + clog "github.com/coredns/coredns/plugin/pkg/log" + pkgparse "github.com/coredns/coredns/plugin/pkg/parse" + "github.com/coredns/coredns/plugin/pkg/transport" +) + +const pluginName = "shed" + +var log = clog.NewWithPlugin(pluginName) + +// stackDepth is a burst budget, not a knob: ~12-16ms at a typical socket's +// serialized drain rate; a larger value would only hold staler responses. +const stackDepth = 1024 + +func init() { plugin.Register(pluginName, setup) } + +func setup(c *caddy.Controller) error { + s, err := parse(c) + if err != nil { + return plugin.Error(pluginName, err) + } + + // The hook only exists on the plain-DNS UDP path; on any other transport + // shed would silently protect nothing, so refuse at parse time. Every + // block key is checked: caddy propagates the plugin list to all keys of + // a server block, not just the one setup runs for. + for _, key := range c.ServerBlockKeys { + if tr, _ := pkgparse.Transport(key); tr != transport.DNS { + return plugin.Error(pluginName, c.Errf("only plain DNS server blocks are supported; %q uses transport %q", key, tr)) + } + } + + cfg := dnsserver.GetConfig(c) + cfg.UDPDecorateWriterFunc = s.decorateWriterFactory + + // On a graceful reload the retiring instance must stop its writer + // goroutines; shutdown removes only this instance's sockets. + c.OnShutdown(s.shutdown) + + cfg.AddPlugin(func(next plugin.Handler) plugin.Handler { + s.Next = next + return s + }) + return nil +} + +func parse(c *caddy.Controller) (*Shed, error) { + s := &Shed{} + i := 0 + for c.Next() { + if i > 0 { + return nil, plugin.ErrOnce + } + i++ + if len(c.RemainingArgs()) != 0 { + return nil, c.ArgErr() + } + if c.NextBlock() { + return nil, c.Errf("shed takes no options") + } + } + return s, nil +} diff --git a/plugin/shed/setup_test.go b/plugin/shed/setup_test.go new file mode 100644 index 000000000..1287ec977 --- /dev/null +++ b/plugin/shed/setup_test.go @@ -0,0 +1,53 @@ +package shed + +import ( + "strings" + "testing" + + "github.com/coredns/caddy" +) + +func TestSetup(t *testing.T) { + tests := []struct { + input string + shouldErr bool + }{ + {"shed", false}, + {"shed extra", true}, + {"shed {\n depth 10\n}", true}, + {"shed\nshed", true}, + } + for i, tc := range tests { + c := caddy.NewTestController("dns", tc.input) + err := setup(c) + if tc.shouldErr && err == nil { + t.Errorf("Test %d: expected error for input %q", i, tc.input) + } + if !tc.shouldErr && err != nil { + t.Errorf("Test %d: unexpected error for input %q: %s", i, tc.input, err) + } + } +} + +func TestSetupRejectsNonDNSTransport(t *testing.T) { + for _, key := range []string{"tls://.:853", "grpc://.:443", "https://.:443", "quic://.:853"} { + c := caddy.NewTestController("dns", "shed") + c.ServerBlockKeys = []string{key} + err := setup(c) + if err == nil { + t.Errorf("expected error for server block key %q", key) + continue + } + if !strings.Contains(err.Error(), "plain DNS") { + t.Errorf("error for %q = %q, want it to mention plain DNS", key, err) + } + } +} + +func TestSetupAcceptsPlainDNSKeys(t *testing.T) { + c := caddy.NewTestController("dns", "shed") + c.ServerBlockKeys = []string{"example.org.:53", "dns://.:53"} + if err := setup(c); err != nil { + t.Fatalf("unexpected error: %s", err) + } +} diff --git a/plugin/shed/shed.go b/plugin/shed/shed.go new file mode 100644 index 000000000..721fcdc9c --- /dev/null +++ b/plugin/shed/shed.go @@ -0,0 +1,116 @@ +// Package shed bounds concurrent UDP response writes per listener socket so +// that an overload storm degrades into counted drops instead of a goroutine +// pile-up and the Go runtime's fdMutex overflow panic. See README.md for the +// failure mode and fdmutex_test.go for its reproduction. +// +// Setup installs dnsserver's Config.UDPDecorateWriterFunc: the decorated +// Write pushes the packed response onto a bounded evict-oldest stack and one +// writer goroutine per socket pops newest-first onto the wire, so the fd +// never sees more than one writer. While a socket's stack is full, ServeDNS +// drops arrivals before any chain work. Every drop is counted. +package shed + +import ( + "context" + "sync" + + "github.com/coredns/coredns/core/dnsserver" + "github.com/coredns/coredns/plugin" + "github.com/coredns/coredns/request" + + "github.com/miekg/dns" + "github.com/prometheus/client_golang/prometheus" +) + +// Shed implements the plugin.Handler interface. +type Shed struct { + Next plugin.Handler +} + +// registry maps a listener socket's *dnsserver.Server to its per-socket +// state. It is package level: several server blocks can share a listener, +// and every block's ServeDNS must see that socket's one stack. +var registry sync.Map // *dnsserver.Server -> *socketState + +// socketState is one listener socket's registry record. owner scopes +// shutdown to this instance's entries on reload. +type socketState struct { + owner *Shed + stack *respStack + droppedQuery prometheus.Counter +} + +// lookupState returns the per-socket state for this request's socket, or nil +// — and on nil ServeDNS fails open. A miss happens when the request carries +// no *dnsserver.Server in its context (tests, non-dnsserver entry points), +// or for a handler still finishing on a server a reload already removed; +// neither admits unbounded new work. +func (s *Shed) lookupState(ctx context.Context) *socketState { + srv := ctx.Value(dnsserver.Key{}) + if srv == nil { + return nil + } + if v, ok := registry.Load(srv); ok { + return v.(*socketState) + } + return nil +} + +// mintState creates one listener socket's state and starts its writer +// goroutine. Idempotent; the loaded path only happens in tests. +func (s *Shed) mintState(srv *dnsserver.Server) *socketState { + st := &socketState{ + owner: s, + stack: newRespStack(stackDepth, droppedTotal.WithLabelValues(srv.Address(), "response")), + droppedQuery: droppedTotal.WithLabelValues(srv.Address(), "query"), + } + if v, loaded := registry.LoadOrStore(srv, st); loaded { + return v.(*socketState) + } + go st.stack.writerLoop() + return st +} + +// shutdown removes this instance's registry entries and stops their writer +// goroutines. A push after removal is rejected by the closed stack and +// counted as a dropped response. +func (s *Shed) shutdown() error { + registry.Range(func(k, v any) bool { + if st := v.(*socketState); st.owner == s { + // LoadAndDelete makes close-once structural even if sweeps race. + if _, loaded := registry.LoadAndDelete(k); loaded { + st.stack.close() + } + } + return true + }) + return nil +} + +// ServeDNS implements the plugin.Handler interface. +func (s *Shed) ServeDNS(ctx context.Context, w dns.ResponseWriter, r *dns.Msg) (int, error) { + // The failure mechanism is UDP-specific: connectionless writes racing + // one fdMutex. TCP must never be starved by UDP-storm shedding. + state := request.Request{W: w, Req: r} + if state.Proto() != "udp" { + return plugin.NextOrFailure(s.Name(), s.Next, ctx, w, r) + } + st := s.lookupState(ctx) + if st == nil { + return plugin.NextOrFailure(s.Name(), s.Next, ctx, w, r) // fail open — see lookupState + } + + // Coupled shed. full() is a racy read by design: a load-shedding + // heuristic, not an invariant. Silent drop: nothing is written, and + // RcodeSuccess satisfies plugin.ClientWrite so dnsserver writes nothing + // either. + if st.stack.full() { + st.droppedQuery.Inc() + return dns.RcodeSuccess, nil + } + + return plugin.NextOrFailure(s.Name(), s.Next, ctx, w, r) +} + +// Name implements the plugin.Handler interface. +func (s *Shed) Name() string { return pluginName } diff --git a/plugin/shed/shed_test.go b/plugin/shed/shed_test.go new file mode 100644 index 000000000..e28fa4544 --- /dev/null +++ b/plugin/shed/shed_test.go @@ -0,0 +1,172 @@ +package shed + +import ( + "context" + "testing" + + "github.com/coredns/coredns/core/dnsserver" + "github.com/coredns/coredns/plugin" + "github.com/coredns/coredns/plugin/pkg/dnstest" + "github.com/coredns/coredns/plugin/test" + + "github.com/miekg/dns" + "github.com/prometheus/client_golang/prometheus/testutil" +) + +// newShed constructs a Shed whose package-level registry entries (and +// writer goroutines) are removed after the test. +func newShed(t *testing.T, next plugin.Handler) *Shed { + t.Helper() + s := &Shed{Next: next} + t.Cleanup(func() { _ = s.shutdown() }) + return s +} + +func msg() *dns.Msg { + m := new(dns.Msg) + m.SetQuestion("example.org.", dns.TypeA) + return m +} + +// packedReply is what miekg/dns's WriteMsg hands the decorated writer. +func packedReply(t *testing.T) []byte { + t.Helper() + m := new(dns.Msg) + m.SetReply(msg()) + data, err := m.Pack() + if err != nil { + t.Fatal(err) + } + return data +} + +func ctxFor(srv *dnsserver.Server) context.Context { + return context.WithValue(context.Background(), dnsserver.Key{}, srv) +} + +// answering is a Next handler that writes a response. +func answering() plugin.Handler { + return plugin.HandlerFunc(func(_ context.Context, w dns.ResponseWriter, r *dns.Msg) (int, error) { + m := new(dns.Msg) + m.SetReply(r) + if err := w.WriteMsg(m); err != nil { + return dns.RcodeServerFailure, err + } + return dns.RcodeSuccess, nil + }) +} + +// blockingWriter parks the writer goroutine in a raw Write until release is +// closed. +type blockingWriter struct { + entered chan struct{} + release chan struct{} +} + +func (b *blockingWriter) Write(p []byte) (int, error) { + b.entered <- struct{}{} + <-b.release + return len(p), nil +} + +// fillStack parks srv's writer goroutine and fills the stack through the +// production decorator; cleanup releases the writer. +func fillStack(t *testing.T, s *Shed, srv *dnsserver.Server) *socketState { + t.Helper() + dec := s.decorateWriterFactory(srv) + v, ok := registry.Load(srv) + if !ok { + t.Fatal("decorator factory must register the socket's state") + } + st := v.(*socketState) + bw := &blockingWriter{ + entered: make(chan struct{}, stackDepth+2), + release: make(chan struct{}), + } + t.Cleanup(func() { close(bw.release) }) + data := packedReply(t) + // First push is popped by the writer, which parks in the raw Write. + if _, err := dec(bw).Write(data); err != nil { + t.Fatal(err) + } + <-bw.entered + for !st.stack.full() { + if _, err := dec(bw).Write(data); err != nil { + t.Fatal(err) + } + } + return st +} + +func TestNoServerInContextFailsOpen(t *testing.T) { + s := newShed(t, answering()) + rec := dnstest.NewRecorder(&test.ResponseWriter{}) + if _, err := s.ServeDNS(context.Background(), rec, msg()); err != nil { + t.Fatal(err) + } + if rec.Msg == nil { + t.Fatal("expected a response without a dnsserver in the context") + } +} + +func TestUnregisteredSocketFailsOpen(t *testing.T) { + s := newShed(t, answering()) + rec := dnstest.NewRecorder(&test.ResponseWriter{}) + // The server carried by the context was never registered by the + // decorator factory (e.g. a straggler after a reload swept it). + if _, err := s.ServeDNS(ctxFor(&dnsserver.Server{}), rec, msg()); err != nil { + t.Fatal(err) + } + if rec.Msg == nil { + t.Fatal("expected a response for an unregistered socket") + } +} + +func TestTCPPassesThrough(t *testing.T) { + s := newShed(t, answering()) + srv := &dnsserver.Server{} + fillStack(t, s, srv) + rec := dnstest.NewRecorder(&test.ResponseWriter{TCP: true}) + // Even with the socket's stack full, TCP is never shed. + if _, err := s.ServeDNS(ctxFor(srv), rec, msg()); err != nil { + t.Fatal(err) + } + if rec.Msg == nil { + t.Fatal("expected a response over TCP") + } +} + +func TestCoupledShedWhenStackFull(t *testing.T) { + s := newShed(t, answering()) + srv := &dnsserver.Server{} + st := fillStack(t, s, srv) + + before := testutil.ToFloat64(st.droppedQuery) + rec := dnstest.NewRecorder(&test.ResponseWriter{}) + rcode, err := s.ServeDNS(ctxFor(srv), rec, msg()) + if err != nil { + t.Fatal(err) + } + if rcode != dns.RcodeSuccess { + t.Errorf("rcode = %d, want RcodeSuccess (silent drop)", rcode) + } + if rec.Msg != nil { + t.Error("a shed query must not be answered") + } + if got := testutil.ToFloat64(st.droppedQuery) - before; got != 1 { + t.Errorf("dropped_total{reason=%q} increment = %v, want 1", "query", got) + } +} + +func TestPassesThroughWhenNotFull(t *testing.T) { + s := newShed(t, answering()) + srv := &dnsserver.Server{} + s.decorateWriterFactory(srv) + rec := dnstest.NewRecorder(&test.ResponseWriter{}) + if _, err := s.ServeDNS(ctxFor(srv), rec, msg()); err != nil { + t.Fatal(err) + } + if rec.Msg == nil { + t.Fatal("expected a response while the stack has room") + } +} diff --git a/plugin/shed/stack.go b/plugin/shed/stack.go new file mode 100644 index 000000000..bddb500d1 --- /dev/null +++ b/plugin/shed/stack.go @@ -0,0 +1,176 @@ +package shed + +import ( + "sync" + "sync/atomic" + + "github.com/coredns/coredns/core/dnsserver" + + "github.com/miekg/dns" + "github.com/prometheus/client_golang/prometheus" +) + +// decorateWriterFactory is installed as Config.UDPDecorateWriterFunc by +// setup. dnsserver's ServePacket calls it once per UDP listener socket, +// before the socket serves its first packet, so the socket's state and +// writer goroutine exist before ServeDNS ever looks them up. miekg/dns then +// applies the returned dns.DecorateWriter once per packet, wrapping the +// response writer. +func (s *Shed) decorateWriterFactory(srv *dnsserver.Server) dns.DecorateWriter { + st := s.mintState(srv) + return func(w dns.Writer) dns.Writer { + return &stackWriter{stack: st.stack, inner: w} + } +} + +// stackWriter is the per-packet transport wrapper. miekg/dns runs every +// message transform (including TSIG) before handing Write the packed bytes, +// so the deferred operation is precisely the serialized syscall. +type stackWriter struct { + stack *respStack + inner dns.Writer // the raw response writer; its Write is the terminal syscall +} + +// Write pushes the packed bytes and reports success: from here on "written" +// means "queued for the socket's writer goroutine". The pushed slice is +// exclusively owned — miekg/dns packs each response into a fresh allocation. +func (sw *stackWriter) Write(data []byte) (int, error) { + if sw.stack.push(pendingResp{w: sw.inner, data: data}) { + sw.stack.dropped.Inc() + } + return len(data), nil +} + +// pendingResp is one captured response awaiting the socket's writer +// goroutine: the packed bytes and the raw writer that puts them on the wire. +type pendingResp struct { + w dns.Writer + data []byte +} + +// respStack is a per-socket bounded ring of pending responses with a single +// insertion cursor and no head index: entries occupy the size slots before +// next (mod depth), so when the ring is full the slot at next holds the +// oldest entry and pushing over it is the eviction. +type respStack struct { + dropped prometheus.Counter // responses evicted, write-failed, or pushed after close + + mu sync.Mutex + buf []pendingResp // ring; len(buf) is the fixed depth + next int // index of the next push + size int // occupied slots + closed bool // set by close; pushes are rejected from then on + + n atomic.Int64 // size mirror for the lock-free full() check + notify chan struct{} // cap 1: writer wake-up + stop chan struct{} // closed on shutdown +} + +func newRespStack(depth int, dropped prometheus.Counter) *respStack { + return &respStack{ + dropped: dropped, + buf: make([]pendingResp, depth), + notify: make(chan struct{}, 1), + stop: make(chan struct{}), + } +} + +// push adds p as the newest entry, evicting the oldest when full. Never +// blocks. Reports whether a response was dropped as a result: the evicted +// oldest, or — on a closed stack — p itself. +func (rs *respStack) push(p pendingResp) (dropped bool) { + rs.mu.Lock() + switch { + case rs.closed: + rs.mu.Unlock() + return true + case rs.size == len(rs.buf): + dropped = true // the slot at next holds the oldest entry + default: + rs.size++ + } + rs.buf[rs.next] = p + rs.next = (rs.next + 1) % len(rs.buf) + rs.n.Store(int64(rs.size)) + rs.mu.Unlock() + select { + case rs.notify <- struct{}{}: + default: + } + return dropped +} + +// pop removes and returns the newest entry. +func (rs *respStack) pop() (pendingResp, bool) { + rs.mu.Lock() + if rs.size == 0 { + rs.mu.Unlock() + return pendingResp{}, false + } + rs.next = (rs.next - 1 + len(rs.buf)) % len(rs.buf) + p := rs.buf[rs.next] + rs.buf[rs.next] = pendingResp{} // release the response bytes + rs.size-- + rs.n.Store(int64(rs.size)) + rs.mu.Unlock() + return p, true +} + +// full is the lock-free view used by the coupled-shed predicate. +func (rs *respStack) full() bool { return rs.n.Load() >= int64(len(rs.buf)) } + +// close stops the writer goroutine — it drains whatever is stacked, then +// exits — and rejects any straggler pushes. +func (rs *respStack) close() { + rs.mu.Lock() + rs.closed = true + rs.mu.Unlock() + close(rs.stop) +} + +// waitNonempty blocks until the stack has work, or reports false once the +// stack is closed and empty. The re-check on the stop arm matters: the +// select may pick stop over a pending notify, but entries accepted before +// the close must still be served — closed guarantees no new pushes, so the +// drain terminates. +func (rs *respStack) waitNonempty() bool { + if rs.n.Load() > 0 { + return true + } + select { + case <-rs.notify: + return true + case <-rs.stop: + return rs.n.Load() > 0 + } +} + +// writerLoop is the socket's single writer: it waits for pending responses, +// pops the one that is newest at write time, and writes it to the wire. +func (rs *respStack) writerLoop() { + for { + if !rs.waitNonempty() { + return + } + if p, ok := rs.pop(); ok { + rs.write(p) + } + } +} + +// write performs the deferred raw write; a response that fails to reach the +// wire is a counted drop. A writer panic must not kill the process — that is +// the failure class this plugin removes — so it is recovered, like +// dnsserver does for synchronous writes. +func (rs *respStack) write(p pendingResp) { + defer func() { + if rec := recover(); rec != nil { + rs.dropped.Inc() + log.Errorf("Recovered panic in shed writer: %v", rec) + } + }() + if _, err := p.w.Write(p.data); err != nil { + rs.dropped.Inc() + log.Debugf("Deferred response write failed: %s", err) + } +} diff --git a/plugin/shed/stack_test.go b/plugin/shed/stack_test.go new file mode 100644 index 000000000..6b4b884a8 --- /dev/null +++ b/plugin/shed/stack_test.go @@ -0,0 +1,164 @@ +package shed + +import ( + "sync/atomic" + "testing" + "time" + + "github.com/coredns/coredns/core/dnsserver" + + "github.com/miekg/dns" +) + +// chanWriter hands each written payload to a channel — the race-safe way to +// observe the writer goroutine's deferred writes. +type chanWriter struct { + got chan []byte +} + +func (w *chanWriter) Write(p []byte) (int, error) { + w.got <- p + return len(p), nil +} + +func TestStackEvictsOldestPopsNewest(t *testing.T) { + rs := newRespStack(3, droppedTotal.WithLabelValues(t.Name(), "response")) // no writer goroutine: pure data structure test + for i := 1; i <= 5; i++ { + dropped := rs.push(pendingResp{data: []byte{byte(i)}}) + if want := i > 3; dropped != want { + t.Errorf("push %d: dropped = %v, want %v", i, dropped, want) + } + } + if !rs.full() { + t.Error("expected full stack after overfilling") + } + // 1 and 2 were evicted; the survivors pop newest-first. + for _, want := range []byte{5, 4, 3} { + p, ok := rs.pop() + if !ok || p.data[0] != want { + t.Fatalf("pop = %v, %v; want entry %d", p.data, ok, want) + } + } + if _, ok := rs.pop(); ok { + t.Error("expected empty stack") + } +} + +func TestStackCloseRejectsPushDrainsRest(t *testing.T) { + rs := newRespStack(4, droppedTotal.WithLabelValues(t.Name(), "response")) + + // A stale notify token on an empty open stack wakes the writer, which + // must tolerate the failed pop (writerLoop's pop-ok check). + rs.push(pendingResp{data: []byte{9}}) + rs.pop() // pop directly, leaving the push's token buffered + if !rs.waitNonempty() { + t.Error("a stale token should report as work") + } + if _, ok := rs.pop(); ok { + t.Error("pop should find nothing behind a stale token") + } + + rs.push(pendingResp{data: []byte{1}}) + rs.close() + if !rs.push(pendingResp{data: []byte{2}}) { + t.Error("push on closed stack should report a drop") + } + // Entries accepted before the close must still be served. + if !rs.waitNonempty() { + t.Fatal("waitNonempty should report the pre-close entry") + } + if p, ok := rs.pop(); !ok || p.data[0] != 1 { + t.Fatalf("pop = %v, %v; want pre-close entry", p, ok) + } + // The pre-close push's token may still be buffered; drain it so the + // final wait deterministically takes the stop arm. + select { + case <-rs.notify: + default: + } + if rs.waitNonempty() { + t.Error("waitNonempty should report false once closed and drained") + } +} + +func TestDecoratorCapturesWriteAndWriterWrites(t *testing.T) { + s := newShed(t, nil) + srv := &dnsserver.Server{} + dec := s.decorateWriterFactory(srv) + if _, ok := registry.Load(srv); !ok { + t.Fatal("factory should pre-register the socket's state") + } + + cw := &chanWriter{got: make(chan []byte, 1)} + data := packedReply(t) + w := dec(cw) + if _, ok := w.(*stackWriter); !ok { + t.Fatalf("decorator returned %T, want *stackWriter", w) + } + if _, err := w.Write(data); err != nil { + t.Fatal(err) + } + select { + case got := <-cw.got: + m := new(dns.Msg) + if err := m.Unpack(got); err != nil { + t.Fatalf("writer goroutine wrote unparseable bytes: %s", err) + } + case <-time.After(5 * time.Second): + t.Fatal("writer goroutine never performed the deferred write") + } +} + +// panicWriter panics on its first Write, then counts. +type panicWriter struct { + writes atomic.Int64 +} + +func (w *panicWriter) Write(p []byte) (int, error) { + if w.writes.Add(1) == 1 { + panic("writer exploded") + } + return len(p), nil +} + +func TestWriterPanicRecovered(t *testing.T) { + s := newShed(t, nil) + srv := &dnsserver.Server{} + dec := s.decorateWriterFactory(srv) + + pw := &panicWriter{} + data := packedReply(t) + if _, err := dec(pw).Write(data); err != nil { + t.Fatal(err) + } + if _, err := dec(pw).Write(data); err != nil { + t.Fatal(err) + } + // The writer goroutine must survive the first write's panic and still + // perform the second. + deadline := time.Now().Add(5 * time.Second) + for pw.writes.Load() < 2 { + if time.Now().After(deadline) { + t.Fatalf("writer performed %d writes, want 2 (goroutine died on panic?)", pw.writes.Load()) + } + time.Sleep(time.Millisecond) + } +} + +func TestShutdownIsInstanceScoped(t *testing.T) { + old := newShed(t, nil) + cur := newShed(t, nil) + oldSrv, newSrv := &dnsserver.Server{}, &dnsserver.Server{} + old.decorateWriterFactory(oldSrv) + cur.decorateWriterFactory(newSrv) + + if err := old.shutdown(); err != nil { + t.Fatal(err) + } + if _, ok := registry.Load(oldSrv); ok { + t.Error("old instance's entry should be swept") + } + if _, ok := registry.Load(newSrv); !ok { + t.Error("new instance's entry must survive the old instance's shutdown") + } +} diff --git a/test/shed_test.go b/test/shed_test.go new file mode 100644 index 000000000..a92c0c3ed --- /dev/null +++ b/test/shed_test.go @@ -0,0 +1,47 @@ +package test + +import ( + "testing" + + "github.com/coredns/coredns/plugin/metrics" + "github.com/coredns/coredns/plugin/test" + + "github.com/miekg/dns" +) + +// TestShed checks that with the shed plugin installed a query is answered +// over both UDP (through the plugin's deferred single-writer path) and +// TCP (which shed passes through), and that its drop counter is exported. +func TestShed(t *testing.T) { + corefile := `.:0 { + shed + prometheus localhost:0 + whoami + }` + + i, udp, tcp, err := CoreDNSServerAndPorts(corefile) + if err != nil { + t.Fatalf("Could not get CoreDNS serving instance: %s", err) + } + defer i.Stop() + + m := new(dns.Msg) + m.SetQuestion("whoami.example.org.", dns.TypeA) + + if r, err := dns.Exchange(m, udp); err != nil || r.Rcode != dns.RcodeSuccess { + t.Fatalf("Expected UDP reply, got %v: %v", r, err) + } + c := &dns.Client{Net: "tcp"} + if r, _, err := c.Exchange(m, tcp); err != nil || r.Rcode != dns.RcodeSuccess { + t.Fatalf("Expected TCP reply, got %v: %v", r, err) + } + + data := test.Scrape("http://" + metrics.ListenAddr + "/metrics") + got, labels := test.MetricValue("coredns_shed_dropped_total", data) + if got != "0" { + t.Errorf("Expected coredns_shed_dropped_total 0, but got %s", got) + } + if labels["reason"] == "" { + t.Errorf("Expected coredns_shed_dropped_total to carry a reason label, got %v", labels) + } +}