diff --git a/cmd/webui4frpc/main.go b/cmd/webui4frpc/main.go index ff6628f..05e13f8 100644 --- a/cmd/webui4frpc/main.go +++ b/cmd/webui4frpc/main.go @@ -175,13 +175,39 @@ func main() { return err } // A re-claim after a forward-centric stop leaves the link flagged - // disabled (by RevokeFn); clear it so renderRemote renders the - // proxy back in. No-op for a fresh claim. - _ = st.SetLinkDisabled(tk.Local.Name, tk.Remote.Name, tk.Link.RemotePort, false) + // disabled (by RevokeFn). Do NOT clear that flag here: the ring + // re-claims automatically (startup reconcile re-spawns any owned + // forward whose worker is missing), so clearing it on claim made + // "stopped" un-durable — every restart resurrected a forward the user + // had explicitly stopped, and it then failed forever against a local + // service that was intentionally not running. + // + // The flag is now owned by the two user-facing actions: + // startForward → SetLinkDisabled(false) before submitting the task + // stopForward → SetLinkDisabled(true) before revoking it + // Claim only READS it to decide whether to bring the worker up. A fresh + // claim of a link with no row still starts, because the lookup miss + // below is treated as "not disabled". + disabled := false + if existing, found, err := st.LinkByTriple(tk.Link.Local, tk.Link.Remote, tk.Link.RemotePort); err != nil { + return err + } else if found { + disabled = existing.Disabled + } // Start (or restart) the per-forward worker for exactly this link. // Each forward has its own frpc process (keyed by the forward // triple); restarting only this key leaves sibling forwards' // processes untouched. + if disabled { + // A stopped forward must also not linger in the topology: leaving + // the entry behind is what let the reconcile loop above keep + // re-claiming it on every restart. + if ring != nil { + ring.RemoveTopologyEntry(tk.Local.Name, tk.Remote.Name, tk.Link.RemotePort) + } + log.Printf("ring[%s] claim %s skipped: %s→%s:%d is disabled", selfID, tk.ID, tk.Local.Name, tk.Remote.Name, tk.Link.RemotePort) + return nil + } if tk.Remote.Enabled { key := process.WorkerKey(tk.Local.Name, tk.Remote.Name, tk.Link.RemotePort) if _, has := pm.Status(key); has { @@ -205,6 +231,14 @@ func main() { if _, running := pm.Status(key); running { _ = pm.Stop(key) } + // Drop the topology entry too. Leaving it behind meant the entry + // outlived the worker, and the next startup reconcile saw a + // "missing" worker for an owned forward and re-claimed it — which + // restarted a forward the user had explicitly stopped. Revoking + // must be a complete retirement, not just a stop. + if ring != nil { + ring.RemoveTopologyEntry(tk.Local.Name, tk.Remote.Name, tk.Link.RemotePort) + } log.Printf("ring[%s] revoked task %s: %s→%s:%d", selfID, tk.ID, tk.Local.Name, tk.Remote.Name, tk.Link.RemotePort) return nil }, @@ -437,6 +471,16 @@ func main() { if t.OwnerID != selfID { continue } + // A forward the user stopped must stay stopped: skip it here so a + // restart does not re-spawn its worker. The ClaimFn enforces the + // same rule (belt and braces — this also avoids a pointless + // claim→skip round trip per disabled forward on every boot). + if ln, found, err := st.LinkByTriple(t.Local.Name, t.Remote.Name, t.Link.RemotePort); err != nil { + log.Printf("ring[%s] reconcile: lookup %s→%s:%d: %v", ring.ID, t.Local.Name, t.Remote.Name, t.Link.RemotePort, err) + continue + } else if found && ln.Disabled { + continue + } key := process.WorkerKey(t.Local.Name, t.Remote.Name, t.Link.RemotePort) if _, has := pm.Status(key); has { continue // worker already running diff --git a/internal/cluster/debug.go b/internal/cluster/debug.go new file mode 100644 index 0000000..900af05 --- /dev/null +++ b/internal/cluster/debug.go @@ -0,0 +1,91 @@ +package cluster + +import ( + "log" + "os" + "sync/atomic" +) + +// logf is the package's logging seam: every cluster log line funnels through it +// so the debug gate in debug.go has a single place to hook. +func logf(format string, args ...any) { log.Printf(format, args...) } + +// Token rotation is the ring's heartbeat: with a 2s round delay it fires +// continuously, and the default three log lines per round +// ("OnToken cycle=N" / "forward cycle=N to X" / "token-send ... -> 200") +// carry no information — no cycle number, address, timing or payload ever +// changes in the steady state. Measured on the live 3-node cluster that was +// ~510k lines/day on one node (3.5M lines in a week), which drowns every real +// event in the journal and fills the disk for a signal that is already +// available in structured form via GET /api/manager/cluster/ring (cycle, +// lastSync, roundDelayMs, node aliveness). +// +// So the steady-state lines are demoted to a debug level, off by default and +// enabled with W4F_DEBUG=token (or 1/all/true for every debug line). Failure +// paths are NOT demoted: a send error, a stale token, a timeout or a leader +// change is exactly what someone is grepping for, and losing those to a quiet +// default would be a bad trade. +const debugEnv = "W4F_DEBUG" + +// debugToken logs a per-token-round heartbeat line. Suppressed unless +// W4F_DEBUG selects "token" (or a catch-all value). +func debugToken(format string, args ...any) { + if debugTokenOn.Load() { + logf(format, args...) + } +} + +// debugAll logs an ad-hoc diagnostic line. Suppressed unless W4F_DEBUG is set +// to a catch-all value (1/all/true/*). +func debugAll(format string, args ...any) { + if debugAllOn.Load() { + logf(format, args...) + } +} + +var ( + debugTokenOn atomic.Bool + debugAllOn atomic.Bool +) + +func init() { ReloadDebug() } + +// ReloadDebug re-reads W4F_DEBUG. Called once at init so tests can flip it +// without restarting, and available at runtime for an operator who wants to +// watch the ring without a redeploy. +func ReloadDebug() { + v := os.Getenv(debugEnv) + switch normalized := normalizeDebugValue(v); normalized { + case "token": + debugTokenOn.Store(true) + debugAllOn.Store(false) + case "all": + debugTokenOn.Store(true) + debugAllOn.Store(true) + default: + debugTokenOn.Store(false) + debugAllOn.Store(false) + } +} + +func normalizeDebugValue(v string) string { + // Compare case-insensitively without pulling in strings just for this. + out := make([]rune, 0, len(v)) + for _, r := range v { + if r >= 'A' && r <= 'Z' { + r += 'a' - 'A' + } + out = append(out, r) + } + s := string(out) + switch s { + case "": + return "" + case "token", "tokens", "ring": + return "token" + } + // Any other non-empty value is a deliberate request for more output, so it + // is treated as a catch-all rather than silently muting the operator who + // set it. "0" lands here too: it was asked for, so honour it. + return "all" +} diff --git a/internal/cluster/debug_test.go b/internal/cluster/debug_test.go new file mode 100644 index 0000000..70efe4c --- /dev/null +++ b/internal/cluster/debug_test.go @@ -0,0 +1,108 @@ +package cluster + +import ( + "bytes" + "log" + "os" + "strings" + "testing" +) + +// captureLog redirects the standard logger into a buffer for the duration of +// fn and returns what was written. +func captureLog(t *testing.T, fn func()) string { + t.Helper() + var buf bytes.Buffer + orig := log.Writer() + origFlags := log.Flags() + log.SetOutput(&buf) + log.SetFlags(0) + defer func() { + log.SetOutput(orig) + log.SetFlags(origFlags) + }() + fn() + return buf.String() +} + +func TestDebugTokenSuppressedByDefault(t *testing.T) { + os.Unsetenv(debugEnv) + ReloadDebug() + + out := captureLog(t, func() { + debugToken("ring[%s] OnToken cycle=%d", "node:7500", 42) + debugToken("ring[%s] forward cycle=%d to %s", "node:7500", 42, "next:7500") + debugToken("token-send %s: size=%dB elapsed=%v -> %d", "next:7500", 5605, "170ms", 200) + }) + if out != "" { + t.Fatalf("steady-state token lines logged with W4F_DEBUG unset: %q", out) + } +} + +func TestDebugTokenEnabledByEnv(t *testing.T) { + for _, v := range []string{"token", "TOKEN", "tokens", "ring", "1", "true", "yes", "all", "*", "yes-please"} { + t.Run(v, func(t *testing.T) { + t.Setenv(debugEnv, v) + ReloadDebug() + if !debugTokenOn.Load() { + t.Fatalf("W4F_DEBUG=%q should enable the token heartbeat", v) + } + out := captureLog(t, func() { + debugToken("ring[%s] OnToken cycle=%d", "node:7500", 7) + }) + if !strings.Contains(out, "OnToken cycle=7") { + t.Fatalf("W4F_DEBUG=%q: expected the line to be logged, got %q", v, out) + } + }) + } +} + +// The whole point of the gate is to shrink the journal, so the failure paths +// must stay loud without any env var — losing a send error to a quiet default +// would be a bad trade. +func TestFailurePathsStayLoud(t *testing.T) { + os.Unsetenv(debugEnv) + ReloadDebug() + + out := captureLog(t, func() { + logf("token-send %s: size=%dB elapsed=%v err=%v", "down:7500", 10, "5ms", "connection refused") + }) + if !strings.Contains(out, "connection refused") { + t.Fatalf("send errors must never be gated: got %q", out) + } +} + +func TestDebugAllOffForTokenOnly(t *testing.T) { + t.Setenv(debugEnv, "token") + ReloadDebug() + if !debugTokenOn.Load() { + t.Fatal("token level should enable debugToken") + } + if debugAllOn.Load() { + t.Fatal("token level must not enable the catch-all debugAll") + } + out := captureLog(t, func() { debugAll("scratch diagnostic") }) + if out != "" { + t.Fatalf("debugAll should stay off at token level, got %q", out) + } +} + +func TestNormalizeDebugValue(t *testing.T) { + cases := map[string]string{ + "": "", + "token": "token", + "TOKEN": "token", + "Ring": "token", + "1": "all", + "true": "all", + "ALL": "all", + "*": "all", + "anything": "all", // unrecognised but deliberate → don't silence it + "0": "all", // 0 is still a deliberate request for output + } + for in, want := range cases { + if got := normalizeDebugValue(in); got != want { + t.Errorf("normalizeDebugValue(%q) = %q, want %q", in, got, want) + } + } +} diff --git a/internal/cluster/ring_engine.go b/internal/cluster/ring_engine.go index d953e77..8e9a5c2 100644 --- a/internal/cluster/ring_engine.go +++ b/internal/cluster/ring_engine.go @@ -182,7 +182,9 @@ func (e *Engine) OnToken(ctx context.Context, tk *Token) (*Token, error) { if tk.SentAt > e.lastTokenAt { e.lastTokenAt = tk.SentAt } - log.Printf("ring[%s] OnToken cycle=%d", e.ID, tk.Cycle) + // Per-round heartbeat — steady state, no information. See debug.go: this + // fired ~2x/second and dominated the journal on every node. + debugToken("ring[%s] OnToken cycle=%d", e.ID, tk.Cycle) // Parallel rhythm timer: operations run while the pace clock ticks. // Delay scales with alive node count (more nodes → lower per-hop delay, @@ -792,6 +794,14 @@ func (e *Engine) RemoveNode(nodeID string) *Task { return e.state.AddRemoveNode(nodeID) } +// RemoveTopologyEntry drops the active topology entry for a forward identified +// by its natural key. Exposed so the app's claim path can retire an entry it +// refuses to serve (see ClaimFn: a disabled forward must stop looking active, +// otherwise the startup reconcile keeps re-claiming it on every restart). +func (e *Engine) RemoveTopologyEntry(local, remote string, port int) bool { + return e.state.RemoveTopology(local, remote, port) +} + // RevokeTask publishes a revocation for an established forward through the // same token channel; the owning node stops the worker and drops topology. func (e *Engine) RevokeTask(local store.Local, remote store.Remote, link store.Link) *Task { diff --git a/internal/cluster/ring_leader.go b/internal/cluster/ring_leader.go index 1c8cde5..4b41505 100644 --- a/internal/cluster/ring_leader.go +++ b/internal/cluster/ring_leader.go @@ -128,7 +128,7 @@ func (e *Engine) forwardToNext(ctx context.Context, tk *Token) error { if e.send == nil { return nil } - log.Printf("ring[%s] forward cycle=%d to %s", e.ID, tk.Cycle, next) + debugToken("ring[%s] forward cycle=%d to %s", e.ID, tk.Cycle, next) err := e.send(ctx, next, tk) if err == nil { if e.state.LeaderID == e.ID { diff --git a/internal/cluster/token_transport.go b/internal/cluster/token_transport.go index 1c2a46f..13fcf17 100644 --- a/internal/cluster/token_transport.go +++ b/internal/cluster/token_transport.go @@ -8,7 +8,6 @@ import ( "context" "encoding/json" "fmt" - "log" "net/http" "time" ) @@ -52,11 +51,16 @@ func (t *TokenTransport) SendTo(getAddr func(nodeID string) string) func(ctx con start := time.Now() resp, err := cli.Do(req) if err != nil { - log.Printf("token-send %s: size=%dB elapsed=%v err=%v", addr, len(body), time.Since(start).Round(time.Millisecond), err) + // NOT demoted: a failed send is the "neighbor offline" signal the + // ring's fault paths are diagnosed from. + logf("token-send %s: size=%dB elapsed=%v err=%v", addr, len(body), time.Since(start).Round(time.Millisecond), err) return err } defer resp.Body.Close() - log.Printf("token-send %s: size=%dB elapsed=%v -> %d", addr, len(body), time.Since(start).Round(time.Millisecond), resp.StatusCode) + // Success path is a per-round heartbeat (size/elapsed/200 repeat + // verbatim every cycle) — debug only. The status-code check below stays + // loud on purpose, so a non-2xx still shows up without the debug flag. + debugToken("token-send %s: size=%dB elapsed=%v -> %d", addr, len(body), time.Since(start).Round(time.Millisecond), resp.StatusCode) if resp.StatusCode >= 400 { return fmt.Errorf("token POST %s -> HTTP %d", url, resp.StatusCode) } diff --git a/internal/httpapi/forward_revoke_test.go b/internal/httpapi/forward_revoke_test.go new file mode 100644 index 0000000..da5c018 --- /dev/null +++ b/internal/httpapi/forward_revoke_test.go @@ -0,0 +1,119 @@ +package httpapi + +import ( + "bytes" + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "path/filepath" + "testing" + + "webui4frpc/internal/cluster" + "webui4frpc/internal/process" + "webui4frpc/internal/store" +) + +// newRingTestHandler builds a Handler WITH a ring engine attached, so the +// paths that publish tasks into the token actually execute. newTestHandler +// leaves Ring nil, which silently skips them — a test built on it can pass +// while the publish side is completely broken. +func newRingTestHandler(t *testing.T) (*Handler, *cluster.Engine, *httptest.Server) { + t.Helper() + dir := t.TempDir() + st, err := store.New(filepath.Join(dir, "test.db")) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = st.Close() }) + + pm := process.NewManager(process.Options{ + ConfigsDir: filepath.Join(dir, "configs"), + LogsDir: filepath.Join(dir, "logs"), + BinaryPath: func() string { return "" }, + Render: func(string) ([]byte, error) { return []byte(`{}`), nil }, + AutoRestart: func(string) bool { return false }, + RestartInterval: func() int { return 5 }, + }) + ring := cluster.NewEngine("n1", "n1:7500", "u", "p", "0.1.0", nil, + &cluster.AppHandler{}, + func(ctx context.Context, next string, tk *cluster.Token) error { return nil }, + "n1:7500", true, "") + h := &Handler{Store: st, Process: pm, WorkDir: dir, User: "admin", Password: "pw", Ring: ring} + mux, err := NewServeMux(h) + if err != nil { + t.Fatal(err) + } + ts := httptest.NewServer(mux) + t.Cleanup(ts.Close) + return h, ring, ts +} + +func saveCanvas(t *testing.T, srv *httptest.Server, body string) { + t.Helper() + req, _ := http.NewRequest(http.MethodPut, srv.URL+"/api/manager/canvas", bytes.NewBufferString(body)) + req.SetBasicAuth("admin", "pw") + resp, err := srv.Client().Do(req) + if err != nil { + t.Fatal(err) + } + resp.Body.Close() + if resp.StatusCode != http.StatusOK { + t.Fatalf("canvas save status = %d", resp.StatusCode) + } +} + +// TestStopForwardPublishesDisabledFlagInRevokeTask is the regression test for +// the actual defect. +// +// stopForward() read the link, called SetLinkDisabled(true), and then handed +// the STALE copy (disabled=false) to RevokeTask. The revoke travels to the +// node that OWNS the forward, and that node's Claim/Revoke path keys off the +// flag — so a stale false meant: +// - the owner could not tell the stop was deliberate, and +// - nothing retired the topology entry, +// +// so the next restart's reconcile re-claimed the forward and spawned a worker +// for something the user had stopped (seen live: endless connection-refused +// against an intentionally-down service). +// +// This asserts the flag ON THE PUBLISHED TASK, which is the value that was +// actually wrong. It cannot be satisfied by the store write alone. +func TestStopForwardPublishesDisabledFlagInRevokeTask(t *testing.T) { + _, ring, ts := newRingTestHandler(t) + + saveCanvas(t, ts, `{ + "locals": [{"name":"svc","ip":"127.0.0.1","port":59999,"protocol":"tcp"}], + "remotes": [{"name":"srv-a","ip":"1.2.3.4","port":7000,"enabled":true}], + "links": [{"local":"svc","remote":"srv-a","remotePort":45999}] + }`) + + // Stop the forward over the API. + b, _ := json.Marshal(stopForwardReq{"svc", "srv-a", 45999}) + req, _ := http.NewRequest(http.MethodPost, ts.URL+"/api/manager/forwards/stop", bytes.NewReader(b)) + req.SetBasicAuth("admin", "pw") + resp, err := ts.Client().Do(req) + if err != nil { + t.Fatal(err) + } + resp.Body.Close() + if resp.StatusCode != http.StatusOK { + t.Fatalf("stop status = %d", resp.StatusCode) + } + + // Find the revoke task that was published into the token. + var revoke *cluster.Task + for _, tk := range ring.State().PendingList() { + if tk.Revoke && tk.Local.Name == "svc" && tk.Link.RemotePort == 45999 { + revoke = tk + break + } + } + if revoke == nil { + t.Fatal("stop did not publish a revoke task for the forward") + } + if !revoke.Link.Disabled { + t.Fatal("the published revoke task carries disabled=false — the owner node cannot tell " + + "this stop was deliberate, which is the bug that let stopped forwards resurrect") + } +} diff --git a/internal/httpapi/forward_stop_test.go b/internal/httpapi/forward_stop_test.go new file mode 100644 index 0000000..d912ac3 --- /dev/null +++ b/internal/httpapi/forward_stop_test.go @@ -0,0 +1,141 @@ +package httpapi + +import ( + "bytes" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" +) + +// stopForwardReq mirrors the forwards start/stop request body. +type stopForwardReq struct { + Local string `json:"local"` + Remote string `json:"remote"` + RemotePort int `json:"remotePort"` +} + +func postForwards(t *testing.T, srv *httptest.Server, action string, body stopForwardReq) int { + t.Helper() + b, _ := json.Marshal(body) + req, _ := http.NewRequest(http.MethodPost, srv.URL+"/api/manager/forwards/"+action, bytes.NewReader(b)) + req.SetBasicAuth("admin", "pw") + req.Header.Set("Content-Type", "application/json") + resp, err := srv.Client().Do(req) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + return resp.StatusCode +} + +// TestStopForwardPersistsDisabledFlag is the regression test for the +// "stopped forwards resurrect on restart" bug. +// +// The original defect: stopForward() read the link, called +// SetLinkDisabled(true), but then handed the STALE (disabled=false) copy to +// RevokeTask — so the disabled flag never reached the node owning the forward, +// and nothing removed the topology entry. On the next restart the startup +// reconcile saw an owned forward with no worker and re-claimed it, spawning a +// worker for a forward the user had deliberately stopped (observed live: +// ~14k "proxy already exists" retries and endless connection-refused against a +// service that was intentionally down). +// +// The contract this pins: after a successful stop, the persisted link MUST be +// disabled — that flag is the single source of truth the claim path consults. +func TestStopForwardPersistsDisabledFlag(t *testing.T) { + h, ts := newTestHandler(t) + + body := `{ + "locals": [{"name":"svc","ip":"127.0.0.1","port":59999,"protocol":"tcp"}], + "remotes": [{"name":"srv-a","ip":"1.2.3.4","port":7000,"enabled":true}], + "links": [{"local":"svc","remote":"srv-a","remotePort":45999}] + }` + req, _ := http.NewRequest(http.MethodPut, ts.URL+"/api/manager/canvas", bytes.NewBufferString(body)) + req.SetBasicAuth("admin", "pw") + resp, err := ts.Client().Do(req) + if err != nil { + t.Fatal(err) + } + resp.Body.Close() + if resp.StatusCode != http.StatusOK { + t.Fatalf("canvas save status = %d", resp.StatusCode) + } + + // The forward starts enabled. + ln, found, err := h.Store.LinkByTriple("svc", "srv-a", 45999) + if err != nil || !found { + t.Fatalf("link not persisted: found=%v err=%v", found, err) + } + if ln.Disabled { + t.Fatal("a freshly saved forward must start enabled") + } + + // Stop it. + if code := postForwards(t, ts, "stop", stopForwardReq{"svc", "srv-a", 45999}); code != http.StatusOK { + t.Fatalf("stop status = %d, want 200", code) + } + + // Persisted flag must now be set — this is what the claim path reads. + ln, found, err = h.Store.LinkByTriple("svc", "srv-a", 45999) + if err != nil { + t.Fatal(err) + } + if !found { + t.Fatal("link vanished after stop; stop must be non-destructive") + } + if !ln.Disabled { + t.Fatal("stop did not persist disabled=true — the startup reconcile would resurrect this forward") + } + + // Start must clear it again (the user-facing re-enable path). + if code := postForwards(t, ts, "start", stopForwardReq{"svc", "srv-a", 45999}); code != http.StatusOK { + t.Fatalf("start status = %d, want 200", code) + } + ln, _, err = h.Store.LinkByTriple("svc", "srv-a", 45999) + if err != nil { + t.Fatal(err) + } + if ln.Disabled { + t.Fatal("start did not clear disabled; a re-enabled forward would stay stopped") + } +} + +// TestStopForwardIsNonDestructive pins the per-forward stop semantics the +// revoke path was specifically rewritten for: stopping one forward must not +// touch a sibling forward that shares the same local or remote. +func TestStopForwardIsNonDestructive(t *testing.T) { + h, ts := newTestHandler(t) + + body := `{ + "locals": [{"name":"svc","ip":"127.0.0.1","port":59999,"protocol":"tcp"}], + "remotes": [{"name":"srv-a","ip":"1.2.3.4","port":7000,"enabled":true}], + "links": [ + {"local":"svc","remote":"srv-a","remotePort":45999}, + {"local":"svc","remote":"srv-a","remotePort":46000} + ] + }` + req, _ := http.NewRequest(http.MethodPut, ts.URL+"/api/manager/canvas", bytes.NewBufferString(body)) + req.SetBasicAuth("admin", "pw") + resp, err := ts.Client().Do(req) + if err != nil { + t.Fatal(err) + } + resp.Body.Close() + + if code := postForwards(t, ts, "stop", stopForwardReq{"svc", "srv-a", 45999}); code != http.StatusOK { + t.Fatalf("stop status = %d", code) + } + + stopped, _, _ := h.Store.LinkByTriple("svc", "srv-a", 45999) + sibling, found, _ := h.Store.LinkByTriple("svc", "srv-a", 46000) + if !found { + t.Fatal("sibling forward was destroyed by stopping its neighbour") + } + if !stopped.Disabled { + t.Error("the stopped forward should be disabled") + } + if sibling.Disabled { + t.Error("the sibling forward must stay enabled — per-forward stop, not per-local/remote") + } +} diff --git a/internal/httpapi/handlers_forwards.go b/internal/httpapi/handlers_forwards.go index c43b069..c2ee0e2 100644 --- a/internal/httpapi/handlers_forwards.go +++ b/internal/httpapi/handlers_forwards.go @@ -109,6 +109,16 @@ func (h *Handler) stopForward(local, remote string, port int) error { } else { _ = h.Store.SetLinkDisabled(local, remote, port, true) } + // The revoke task travels to whichever node OWNS the forward, and that node + // re-reads the disabled flag from its own store before starting a worker — + // so the flag has to be set on every node that has a copy of this link, not + // just the one handling this request. Propagating Disabled on the task lets + // the owner's RevokeFn stop the worker even if its own store row is stale. + // + // This also fixes a latent inconsistency: `ln` was read BEFORE the + // SetLinkDisabled(true) above, so the link published into the token still + // carried disabled=false and got copied into the topology entry verbatim. + ln.Disabled = true if loc.LocalOnly { if h.Process != nil { key := process.WorkerKey(local, remote, port) diff --git a/internal/store/link_test.go b/internal/store/link_test.go new file mode 100644 index 0000000..f57bf08 --- /dev/null +++ b/internal/store/link_test.go @@ -0,0 +1,198 @@ +package store + +import ( + "path/filepath" + "testing" +) + +// seed inserts the local + remote rows a link's foreign keys require. +// (links.local / links.remote reference their own tables, so a link cannot +// exist on its own — the same reason ClaimFn upserts them before ReplaceLinks.) +func seed(t *testing.T, st *Store, locals []string, remote string) { + t.Helper() + for _, n := range locals { + if err := st.UpsertLocal(Local{Name: n, IP: "127.0.0.1", Port: 8080, Protocol: "tcp"}); err != nil { + t.Fatal(err) + } + } + if err := st.UpsertRemote(Remote{Name: remote, IP: "1.2.3.4", Port: 7000, Token: "tok", Enabled: true}); err != nil { + t.Fatal(err) + } +} + +// TestLinkByTripleSurvivesReplaceLinks pins the reason LinkByTriple exists. +// +// ReplaceLinks() rewrites the whole table with DELETE + re-INSERT, so sqlite +// hands every row a FRESH autoincrement id. A Link captured before such a +// write (e.g. one riding inside a ring token) therefore carries an id that +// either matches a different forward or matches nothing. The natural key +// (local, remote, remotePort) is what every caller actually identifies a +// forward by, and it must survive those rewrites. +func TestLinkByTripleSurvivesReplaceLinks(t *testing.T) { + st, err := New(filepath.Join(t.TempDir(), "test.db")) + if err != nil { + t.Fatal(err) + } + defer st.Close() + + seed(t, st, []string{"alpha", "beta", "gamma"}, "srv") + links := []Link{ + {Local: "alpha", Remote: "srv", RemotePort: 100}, + {Local: "beta", Remote: "srv", RemotePort: 200}, + {Local: "gamma", Remote: "srv", RemotePort: 300}, + } + if err := st.ReplaceLinks(links); err != nil { + t.Fatal(err) + } + + // Capture the ids as the ring would have them. + before := map[string]int64{} + all, err := st.ListLinks() + if err != nil { + t.Fatal(err) + } + for _, l := range all { + before[l.Local] = l.ID + } + if len(before) != 3 { + t.Fatalf("expected 3 links, got %d", len(before)) + } + + // Rewrite the table (this is what saveCanvas and ClaimFn both do). + if err := st.ReplaceLinks(links); err != nil { + t.Fatal(err) + } + after, err := st.ListLinks() + if err != nil { + t.Fatal(err) + } + if len(after) != 3 { + t.Fatalf("expected 3 links after rewrite, got %d", len(after)) + } + + // The natural key must still resolve to the right forward, with its + // disabled flag and group intact. + for _, l := range after { + if l.Disabled { + t.Errorf("link %s unexpectedly disabled after a plain rewrite", l.Local) + } + } + got, found, err := st.LinkByTriple("beta", "srv", 200) + if err != nil { + t.Fatal(err) + } + if !found { + t.Fatal("LinkByTriple failed to find beta after ReplaceLinks") + } + if got.Local != "beta" || got.RemotePort != 200 { + t.Fatalf("LinkByTriple returned the wrong row: %+v", got) + } +} + +// TestLinkByTripleNotFoundIsNotError documents the contract callers rely on: +// "no persisted opinion yet" is (Link{}, false, nil), not an error. A fresh +// claim of a link with no row must be allowed to start. +func TestLinkByTripleNotFoundIsNotError(t *testing.T) { + st, err := New(filepath.Join(t.TempDir(), "test.db")) + if err != nil { + t.Fatal(err) + } + defer st.Close() + + got, found, err := st.LinkByTriple("nope", "srv", 1234) + if err != nil { + t.Fatalf("missing link must not be an error, got %v", err) + } + if found { + t.Fatalf("missing link reported as found: %+v", got) + } + if got.Local != "" || got.RemotePort != 0 { + t.Fatalf("expected zero Link on miss, got %+v", got) + } +} + +// TestLinkByTripleReadsDisabledFlag is the store-level half of the +// "stopped forwards resurrect on restart" bug: the claim path asks the store +// whether the user disabled this forward, so this lookup must return the flag +// as persisted. +func TestLinkByTripleReadsDisabledFlag(t *testing.T) { + st, err := New(filepath.Join(t.TempDir(), "test.db")) + if err != nil { + t.Fatal(err) + } + defer st.Close() + + seed(t, st, []string{"mc"}, "srv") + if err := st.ReplaceLinks([]Link{{Local: "mc", Remote: "srv", RemotePort: 25565, Group: "game"}}); err != nil { + t.Fatal(err) + } + if err := st.SetLinkDisabled("mc", "srv", 25565, true); err != nil { + t.Fatal(err) + } + + ln, found, err := st.LinkByTriple("mc", "srv", 25565) + if err != nil { + t.Fatal(err) + } + if !found { + t.Fatal("expected to find the link") + } + if !ln.Disabled { + t.Fatal("expected Disabled=true to be visible through LinkByTriple") + } + if ln.Group != "game" { + t.Fatalf("group should survive, got %q", ln.Group) + } + + // ...and the user-facing start path must be able to clear it again. + if err := st.SetLinkDisabled("mc", "srv", 25565, false); err != nil { + t.Fatal(err) + } + if ln, _, _ := st.LinkByTriple("mc", "srv", 25565); ln.Disabled { + t.Fatal("expected Disabled=false after clearing") + } +} + +// TestGetLinkByIDIsStaleAfterReplaceLinks documents WHY callers must not use +// GetLink(id) with a previously captured id. It is not a fix — it is the trap +// being pinned shut, so the hazard stays visible if someone reintroduces it. +func TestGetLinkByIDIsStaleAfterReplaceLinks(t *testing.T) { + st, err := New(filepath.Join(t.TempDir(), "test.db")) + if err != nil { + t.Fatal(err) + } + defer st.Close() + + seed(t, st, []string{"alpha", "beta", "gamma"}, "srv") + if err := st.ReplaceLinks([]Link{ + {Local: "alpha", Remote: "srv", RemotePort: 100}, + {Local: "beta", Remote: "srv", RemotePort: 200}, + }); err != nil { + t.Fatal(err) + } + all, _ := st.ListLinks() + var staleID int64 + for _, l := range all { + if l.Local == "alpha" { + staleID = l.ID + } + } + + if err := st.ReplaceLinks([]Link{ + {Local: "alpha", Remote: "srv", RemotePort: 100}, + {Local: "beta", Remote: "srv", RemotePort: 200}, + {Local: "gamma", Remote: "srv", RemotePort: 300}, + }); err != nil { + t.Fatal(err) + } + + // The old id may still resolve, but to whatever row now occupies that + // id — which is exactly the silent-mis-target hazard. Assert that the + // natural key remains the only safe handle. + if ln, ok := st.GetLink(staleID); ok && ln.Local != "alpha" { + t.Logf("stale id %d now points at %q (hazard confirmed; use LinkByTriple)", staleID, ln.Local) + } + if got, found, _ := st.LinkByTriple("alpha", "srv", 100); !found || got.Local != "alpha" { + t.Fatalf("natural key must stay reliable, got %+v found=%v", got, found) + } +} diff --git a/internal/store/store.go b/internal/store/store.go index f67aa97..39633da 100644 --- a/internal/store/store.go +++ b/internal/store/store.go @@ -560,6 +560,30 @@ func (s *Store) GetLink(id int64) (Link, bool) { return l, true } +// LinkByTriple looks a link up by its natural key (local, remote, remotePort). +// +// Prefer this over GetLink(id) whenever the caller only knows the forward's +// identity: ReplaceLinks() rewrites the whole table with DELETE + re-INSERT, so +// every row gets a fresh autoincrement id. Any id captured before such a write +// (e.g. a Link carried inside a ring token) is stale by definition and will +// either miss or — worse — match a different forward. The natural key is +// stable across those rewrites. +// +// Returns (link, found). A missing row is (Link{}, false) and is NOT an error: +// callers use that to mean "no persisted opinion yet". +func (s *Store) LinkByTriple(local, remote string, port int) (Link, bool, error) { + var l Link + err := s.db.QueryRow("SELECT id, local, remote, remote_port, offset_x, offset_y, grp, disabled FROM links WHERE local = ? AND remote = ? AND remote_port = ?", local, remote, port). + Scan(&l.ID, &l.Local, &l.Remote, &l.RemotePort, &l.OffsetX, &l.OffsetY, &l.Group, &l.Disabled) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return Link{}, false, nil + } + return Link{}, false, err + } + return l, true, nil +} + // DeleteLink removes a single link by id. func (s *Store) DeleteLink(id int64) error { _, err := s.db.Exec("DELETE FROM links WHERE id = ?", id)