diff --git a/records.go b/records.go index 58bd8fc..93178cb 100644 --- a/records.go +++ b/records.go @@ -49,8 +49,11 @@ type aEntry struct { } // nEntry wraps an N record with creation time for TTL expiry. +// Owner is the node ID that registered the record, or 0 when the +// registrant could not be identified. type nEntry struct { NetID uint16 + Owner uint32 CreatedAt time.Time } @@ -63,6 +66,13 @@ type RecordStore struct { storePath string // path to persist records (empty = no persistence) ttl time.Duration done chan struct{} + saveSeq uint64 // guarded by mu; incremented once per snapshot + + // writeMu serializes on-disk writes and guards writtenSeq. It is + // acquired only after mu has been released, so a file write never + // blocks readers of the in-memory maps. + writeMu sync.Mutex + writtenSeq uint64 } type svcKey struct { @@ -115,7 +125,6 @@ func (rs *RecordStore) reapLoop() { func (rs *RecordStore) reapExpired() { rs.mu.Lock() - defer rs.mu.Unlock() now := time.Now() reaped := false @@ -147,9 +156,13 @@ func (rs *RecordStore) reapExpired() { rs.sRecords[key] = alive } } + var pending *pendingSave if reaped { - rs.save() + pending = rs.snapshotLocked() } + rs.mu.Unlock() + + rs.writeSnapshot(pending) } // SetTTL overrides the default record TTL. @@ -189,6 +202,7 @@ type snapshotA struct { type snapshotN struct { Name string `json:"name"` NetworkID uint16 `json:"network_id"` + NodeID uint32 `json:"node_id,omitempty"` CreatedAt time.Time `json:"created_at,omitempty"` } @@ -200,9 +214,23 @@ type snapshotS struct { CreatedAt time.Time `json:"created_at,omitempty"` } -func (rs *RecordStore) save() { +// pendingSave is a serialized store snapshot waiting to be written to +// disk. It carries everything writeSnapshot needs so that no store field +// has to be read while the write is in flight. +type pendingSave struct { + path string + data []byte + seq uint64 + aRecords int + nRecords int +} + +// snapshotLocked serializes the current records. The caller must hold mu +// (at least for reading). It returns nil when persistence is disabled. +// The returned snapshot is written by writeSnapshot after mu is released. +func (rs *RecordStore) snapshotLocked() *pendingSave { if rs.storePath == "" { - return + return nil } snap := recordSnapshot{} @@ -210,7 +238,7 @@ func (rs *RecordStore) save() { snap.ARecords = append(snap.ARecords, snapshotA{Name: name, Address: e.Addr.String(), CreatedAt: e.CreatedAt}) } for name, e := range rs.nRecords { - snap.NRecords = append(snap.NRecords, snapshotN{Name: name, NetworkID: e.NetID, CreatedAt: e.CreatedAt}) + snap.NRecords = append(snap.NRecords, snapshotN{Name: name, NetworkID: e.NetID, NodeID: e.Owner, CreatedAt: e.CreatedAt}) } for key, entries := range rs.sRecords { for _, e := range entries { @@ -225,21 +253,52 @@ func (rs *RecordStore) save() { } // MarshalIndent on recordSnapshot is infallible: every field is a - // primitive (string/uint16) or time.Time, no exotic types. + // primitive (string/uint16/uint32) or time.Time, no exotic types. // The error branch is unreachable. data, _ := json.MarshalIndent(snap, "", " ") - dir := filepath.Dir(rs.storePath) + rs.saveSeq++ + return &pendingSave{ + path: rs.storePath, + data: data, + seq: rs.saveSeq, + aRecords: len(rs.aRecords), + nRecords: len(rs.nRecords), + } +} + +// writeSnapshot persists a snapshot produced by snapshotLocked. It must +// be called with mu released: the directory creation and the atomic file +// write (which fsyncs) can take milliseconds, and holding mu across them +// would stall every concurrent lookup. +// +// Concurrent writers are serialized on writeMu, and a snapshot older than +// the last one written is dropped so the file always ends up holding the +// most recent state. +func (rs *RecordStore) writeSnapshot(p *pendingSave) { + if p == nil || p.path == "" { + return + } + + rs.writeMu.Lock() + defer rs.writeMu.Unlock() + + if p.seq <= rs.writtenSeq { + return + } + + dir := filepath.Dir(p.path) if err := os.MkdirAll(dir, 0700); err != nil { slog.Error("create nameserver state directory", "dir", dir, "err", err) return } - if err := fsutil.AtomicWrite(rs.storePath, data); err != nil { + if err := fsutil.AtomicWrite(p.path, p.data); err != nil { slog.Error("write nameserver state", "err", err) return } - slog.Debug("nameserver state saved", "a_records", len(rs.aRecords), "n_records", len(rs.nRecords)) + rs.writtenSeq = p.seq + slog.Debug("nameserver state saved", "a_records", p.aRecords, "n_records", p.nRecords) } func (rs *RecordStore) load() { @@ -278,7 +337,7 @@ func (rs *RecordStore) load() { rs.aRecords[normalizeName(a.Name)] = &aEntry{Addr: addr, CreatedAt: restore(a.CreatedAt)} } for _, n := range snap.NRecords { - rs.nRecords[normalizeName(n.Name)] = &nEntry{NetID: n.NetworkID, CreatedAt: restore(n.CreatedAt)} + rs.nRecords[normalizeName(n.Name)] = &nEntry{NetID: n.NetworkID, Owner: n.NodeID, CreatedAt: restore(n.CreatedAt)} } for _, s := range snap.SRecords { addr, err := protocol.ParseAddr(s.Address) @@ -310,9 +369,11 @@ func normalizeName(name string) string { func (rs *RecordStore) RegisterA(name string, addr protocol.Addr) { name = normalizeName(name) rs.mu.Lock() - defer rs.mu.Unlock() rs.aRecords[name] = &aEntry{Addr: addr, CreatedAt: time.Now()} - rs.save() + pending := rs.snapshotLocked() + rs.mu.Unlock() + + rs.writeSnapshot(pending) } // LookupA resolves a name to an address. @@ -327,13 +388,32 @@ func (rs *RecordStore) LookupA(name string) (protocol.Addr, error) { return e.Addr, nil } -// RegisterN adds a network name record. +// RegisterN adds a network name record with no recorded owner. func (rs *RecordStore) RegisterN(name string, netID uint16) { + rs.RegisterNOwned(name, netID, 0, false) +} + +// RegisterNOwned adds a network name record and records owner as the +// node that registered it (0 means the registrant is unknown). +// +// When enforce is true and the name already carries a different non-zero +// owner, the record is left untouched and false is returned. With enforce +// false the record is overwritten unconditionally and true is returned. +func (rs *RecordStore) RegisterNOwned(name string, netID uint16, owner uint32, enforce bool) bool { name = normalizeName(name) rs.mu.Lock() - defer rs.mu.Unlock() - rs.nRecords[name] = &nEntry{NetID: netID, CreatedAt: time.Now()} - rs.save() + if enforce { + if cur, ok := rs.nRecords[name]; ok && cur.Owner != 0 && cur.Owner != owner { + rs.mu.Unlock() + return false + } + } + rs.nRecords[name] = &nEntry{NetID: netID, Owner: owner, CreatedAt: time.Now()} + pending := rs.snapshotLocked() + rs.mu.Unlock() + + rs.writeSnapshot(pending) + return true } // LookupN resolves a network name to a network ID. @@ -352,7 +432,6 @@ func (rs *RecordStore) LookupN(name string) (uint16, error) { func (rs *RecordStore) RegisterS(name string, addr protocol.Addr, networkID, port uint16) { name = normalizeName(name) rs.mu.Lock() - defer rs.mu.Unlock() key := svcKey{NetworkID: networkID, Port: port} entry := ServiceEntry{Name: name, Address: addr, Port: port, CreatedAt: time.Now()} // Avoid duplicates — refresh TTL if already present. Also update @@ -364,11 +443,15 @@ func (rs *RecordStore) RegisterS(name string, addr protocol.Addr, networkID, por if e.Address == addr && e.Port == port { rs.sRecords[key][i].Name = name rs.sRecords[key][i].CreatedAt = time.Now() + rs.mu.Unlock() return } } rs.sRecords[key] = append(rs.sRecords[key], entry) - rs.save() + pending := rs.snapshotLocked() + rs.mu.Unlock() + + rs.writeSnapshot(pending) } // LookupS finds service providers on a network+port. @@ -385,9 +468,11 @@ func (rs *RecordStore) LookupS(networkID, port uint16) []ServiceEntry { func (rs *RecordStore) UnregisterA(name string) { name = normalizeName(name) rs.mu.Lock() - defer rs.mu.Unlock() delete(rs.aRecords, name) - rs.save() + pending := rs.snapshotLocked() + rs.mu.Unlock() + + rs.writeSnapshot(pending) } // AllA returns all A records. diff --git a/server.go b/server.go index 901b280..6ea6115 100644 --- a/server.go +++ b/server.go @@ -6,11 +6,28 @@ import ( "fmt" "log/slog" "net" + "os" + "strings" + "sync/atomic" "github.com/pilot-protocol/common/coreapi" "github.com/pilot-protocol/common/protocol" ) +// StrictRegisterEnv names the environment variable that turns on strict +// REGISTER handling. It is off unless the value is "1" or "true". +// +// With strict handling on: +// - every REGISTER must arrive on a connection whose remote address +// resolves to a non-zero node ID; +// - N records are bound to the node that registered them, and a later +// REGISTER N for the same name from a different node is refused. +// +// With it off (the default) the server keeps its previous behaviour: +// unresolvable callers are accepted and N records may be overwritten by +// any caller. +const StrictRegisterEnv = "PILOT_NAMESERVER_STRICT_REGISTER" + // PortListener abstracts the ability to listen on a Pilot overlay port. // Satisfied by *driver.Driver (via a thin wrapper in cmd/nameserver). type PortListener interface { @@ -29,6 +46,7 @@ type Server struct { listener PortListener ln net.Listener ready chan struct{} + strict atomic.Bool } // New creates a nameserver backed by a fresh record store. @@ -38,13 +56,33 @@ func New(pl PortListener, storePath string) *Server { if storePath != "" { store.SetStorePath(storePath) } - return &Server{ + s := &Server{ store: store, listener: pl, ready: make(chan struct{}), } + s.strict.Store(strictRegisterFromEnv()) + return s +} + +// strictRegisterFromEnv reads StrictRegisterEnv. Anything other than +// "1"/"true" (case-insensitive) leaves strict handling off. +func strictRegisterFromEnv() bool { + switch strings.ToLower(strings.TrimSpace(os.Getenv(StrictRegisterEnv))) { + case "1", "true": + return true + default: + return false + } } +// SetStrictRegister turns strict REGISTER handling on or off at runtime, +// overriding whatever StrictRegisterEnv selected at construction. +func (s *Server) SetStrictRegister(on bool) { s.strict.Store(on) } + +// StrictRegister reports whether strict REGISTER handling is on. +func (s *Server) StrictRegister() bool { return s.strict.Load() } + // Ready returns a channel that is closed once the server is listening. func (s *Server) Ready() <-chan struct{} { return s.ready @@ -149,8 +187,17 @@ func (s *Server) handleRegister(req Request, remoteAddr net.Addr) string { return FormatResponseErr(fmt.Sprintf("name too long: %d bytes (max %d)", len(req.Name), MaxNameLength)) } - // Extract caller's node ID from RemoteAddr for source validation + strict := s.strict.Load() + + // Extract caller's node ID from RemoteAddr for source validation. callerNode := extractCallerNode(remoteAddr) + if strict { + node, ok := resolveCallerNode(remoteAddr) + if !ok { + return FormatResponseErr("caller node identity required") + } + callerNode = node + } switch req.RecordType { case RecordA: @@ -167,8 +214,12 @@ func (s *Server) handleRegister(req Request, remoteAddr net.Addr) string { return FormatResponseOK() case RecordN: - s.store.RegisterN(req.Name, req.NetID) - slog.Debug("nameserver registered N record", "name", req.Name, "network_id", req.NetID) + // The owning node is always recorded; it is only enforced against + // a later registrant when strict handling is on. + if !s.store.RegisterNOwned(req.Name, req.NetID, callerNode, strict) { + return FormatResponseErr("network name registered by another node") + } + slog.Debug("nameserver registered N record", "name", req.Name, "network_id", req.NetID, "node", callerNode) return FormatResponseOK() case RecordS: @@ -208,3 +259,45 @@ func extractCallerNode(addr net.Addr) uint32 { } return pilotAddr.Node } + +// resolveCallerNode returns the node ID carried by a remote address, and +// whether one could be determined at all. +// +// A Pilot address renders as ":..", so the +// "
:" string reported by the overlay connection adapter +// contains two colons and is not a host:port pair net.SplitHostPort can +// split. Each plausible substring is tried in turn: the whole string, the +// host part when the string really is host:port (the bracketed form), and +// the string with a trailing ":" removed. +// +// A node ID of 0 is reported as unresolved: it is the zero value used +// throughout this package to mean "no caller identity". +func resolveCallerNode(addr net.Addr) (uint32, bool) { + if addr == nil { + return 0, false + } + s := strings.TrimSpace(addr.String()) + if s == "" { + return 0, false + } + + candidates := []string{s} + if host, _, err := net.SplitHostPort(s); err == nil && host != "" { + candidates = append(candidates, host) + } + if i := strings.LastIndex(s, ":"); i > 0 { + candidates = append(candidates, s[:i]) + } + + for _, c := range candidates { + pilotAddr, err := protocol.ParseAddr(c) + if err != nil { + continue + } + if pilotAddr.Node == 0 { + return 0, false + } + return pilotAddr.Node, true + } + return 0, false +} diff --git a/zz_persist_ownership_test.go b/zz_persist_ownership_test.go new file mode 100644 index 0000000..97cf62d --- /dev/null +++ b/zz_persist_ownership_test.go @@ -0,0 +1,340 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +package nameserver + +import ( + "fmt" + "os" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/pilot-protocol/common/protocol" +) + +// --------------------------------------------------------------------------- +// Persistence happens outside the store lock (records.go). +// --------------------------------------------------------------------------- + +// TestWriteSnapshotRunsWithoutStoreLock pins that a file write in flight +// does not stall readers: writeMu is held for the duration of the write, +// but mu is released before it starts, so lookups (and the newly written +// record) are visible while the write is still pending. +func TestWriteSnapshotRunsWithoutStoreLock(t *testing.T) { + t.Parallel() + dir := t.TempDir() + rs := NewRecordStore() + defer rs.Close() + rs.SetStorePath(filepath.Join(dir, "ns.json")) + + rs.RegisterA("seed", protocol.Addr{Node: 1}) + + // Stand in for a slow disk: the next write blocks until we release it. + rs.writeMu.Lock() + + done := make(chan struct{}) + go func() { + defer close(done) + rs.RegisterA("pending", protocol.Addr{Node: 2}) + }() + + // The map mutation must land (and lookups must keep working) even + // though the write has not completed. + deadline := time.Now().Add(2 * time.Second) + for { + if _, err := rs.LookupA("pending"); err == nil { + break + } + if time.Now().After(deadline) { + rs.writeMu.Unlock() + t.Fatal("lookup of pending record blocked behind the file write") + } + time.Sleep(time.Millisecond) + } + if _, err := rs.LookupA("seed"); err != nil { + rs.writeMu.Unlock() + t.Fatalf("LookupA(seed) during pending write: %v", err) + } + + select { + case <-done: + t.Fatal("RegisterA returned while the file write was still blocked") + default: + } + + rs.writeMu.Unlock() + select { + case <-done: + case <-time.After(2 * time.Second): + t.Fatal("RegisterA did not finish after the write was released") + } +} + +// TestWriteSnapshotDropsStaleSequence pins that an older snapshot landing +// after a newer one does not roll the file back. +func TestWriteSnapshotDropsStaleSequence(t *testing.T) { + t.Parallel() + dir := t.TempDir() + path := filepath.Join(dir, "ns.json") + rs := NewRecordStore() + defer rs.Close() + rs.SetStorePath(path) + + rs.mu.Lock() + rs.aRecords["older"] = &aEntry{Addr: protocol.Addr{Node: 1}, CreatedAt: time.Now()} + older := rs.snapshotLocked() + rs.mu.Unlock() + + rs.mu.Lock() + rs.aRecords["newer"] = &aEntry{Addr: protocol.Addr{Node: 2}, CreatedAt: time.Now()} + newer := rs.snapshotLocked() + rs.mu.Unlock() + + rs.writeSnapshot(newer) + rs.writeSnapshot(older) + + data, err := os.ReadFile(path) + if err != nil { + t.Fatalf("read store: %v", err) + } + if !strings.Contains(string(data), `"newer"`) { + t.Errorf("stale snapshot overwrote the newer one: %s", data) + } +} + +// TestConcurrentRegistersConvergeOnDisk exercises the snapshot/write split +// under concurrency: the last write must reflect every registration. +func TestConcurrentRegistersConvergeOnDisk(t *testing.T) { + t.Parallel() + dir := t.TempDir() + path := filepath.Join(dir, "ns.json") + rs := NewRecordStore() + defer rs.Close() + rs.SetStorePath(path) + + const n = 32 + done := make(chan struct{}, n) + for i := 0; i < n; i++ { + go func(i int) { + defer func() { done <- struct{}{} }() + rs.RegisterA(fmt.Sprintf("host-%d", i), protocol.Addr{Node: uint32(i + 1)}) + _, _ = rs.LookupA("host-0") + }(i) + } + for i := 0; i < n; i++ { + <-done + } + + // Force one final write so the file reflects the settled state + // regardless of which concurrent snapshot won the race. + rs.RegisterA("final", protocol.Addr{Node: 0xFFFF}) + + data, err := os.ReadFile(path) + if err != nil { + t.Fatalf("read store: %v", err) + } + for i := 0; i < n; i++ { + if !strings.Contains(string(data), fmt.Sprintf(`"host-%d"`, i)) { + t.Fatalf("host-%d missing from persisted snapshot", i) + } + } +} + +// --------------------------------------------------------------------------- +// Caller identity + N-record ownership (server.go). +// --------------------------------------------------------------------------- + +func TestResolveCallerNode_Forms(t *testing.T) { + t.Parallel() + want := protocol.Addr{Network: 0, Node: 0x12345678} + + // The shape the overlay connection adapter reports: a Pilot address + // followed by ":". + if got, ok := resolveCallerNode(stringAddr{s: want.String() + ":53"}); !ok || got != want.Node { + t.Errorf("addr:port form: got 0x%x ok=%v, want 0x%x true", got, ok, want.Node) + } + // Bare address, no port. + if got, ok := resolveCallerNode(stringAddr{s: want.String()}); !ok || got != want.Node { + t.Errorf("bare form: got 0x%x ok=%v, want 0x%x true", got, ok, want.Node) + } + // Bracketed host:port. + if got, ok := resolveCallerNode(stringAddr{s: "[" + want.String() + "]:8080"}); !ok || got != want.Node { + t.Errorf("bracketed form: got 0x%x ok=%v, want 0x%x true", got, ok, want.Node) + } + // Node 0 counts as no identity. + zero := protocol.Addr{Network: 0, Node: 0} + if got, ok := resolveCallerNode(stringAddr{s: zero.String() + ":53"}); ok || got != 0 { + t.Errorf("zero node: got 0x%x ok=%v, want 0 false", got, ok) + } + for _, s := range []string{"", "not-a-pilot-addr", "127.0.0.1:8080"} { + if got, ok := resolveCallerNode(stringAddr{s: s}); ok || got != 0 { + t.Errorf("%q: got 0x%x ok=%v, want 0 false", s, got, ok) + } + } + if got, ok := resolveCallerNode(nil); ok || got != 0 { + t.Errorf("nil addr: got 0x%x ok=%v, want 0 false", got, ok) + } +} + +func callerAddr(node uint32) stringAddr { + return stringAddr{s: protocol.Addr{Network: 0, Node: node}.String() + ":53"} +} + +// TestRegisterN_DefaultAllowsOverwrite pins the default (flag off) +// behaviour: any caller may replace an N record. +func TestRegisterN_DefaultAllowsOverwrite(t *testing.T) { + t.Parallel() + s := New(&fakeListener{}, "") + if s.StrictRegister() { + t.Fatal("strict register should be off by default") + } + + req := Request{Command: "REGISTER", RecordType: RecordN, Name: "shared", NetID: 5} + if got := s.handleRequest(req, callerAddr(1)); got != "OK" { + t.Fatalf("first REGISTER N: got %q", got) + } + req.NetID = 9 + if got := s.handleRequest(req, callerAddr(2)); got != "OK" { + t.Fatalf("overwrite from another node: got %q", got) + } + if got, err := s.store.LookupN("shared"); err != nil || got != 9 { + t.Errorf("LookupN: got (%d, %v), want (9, nil)", got, err) + } +} + +// TestRegisterN_StrictEnforcesOwnership pins that with strict handling on, +// only the node that registered a network name may change it. +func TestRegisterN_StrictEnforcesOwnership(t *testing.T) { + t.Parallel() + s := New(&fakeListener{}, "") + s.SetStrictRegister(true) + + req := Request{Command: "REGISTER", RecordType: RecordN, Name: "owned", NetID: 5} + if got := s.handleRequest(req, callerAddr(1)); got != "OK" { + t.Fatalf("first REGISTER N: got %q", got) + } + + other := Request{Command: "REGISTER", RecordType: RecordN, Name: "owned", NetID: 9} + got := s.handleRequest(other, callerAddr(2)) + if !strings.HasPrefix(got, "ERR ") { + t.Fatalf("REGISTER N from another node: got %q, want ERR", got) + } + if netID, err := s.store.LookupN("owned"); err != nil || netID != 5 { + t.Errorf("record changed despite refusal: got (%d, %v)", netID, err) + } + + // The owner may still update its own record. + req.NetID = 7 + if got := s.handleRequest(req, callerAddr(1)); got != "OK" { + t.Fatalf("owner update: got %q", got) + } + if netID, err := s.store.LookupN("owned"); err != nil || netID != 7 { + t.Errorf("owner update not applied: got (%d, %v)", netID, err) + } +} + +// TestRegisterStrictRequiresCallerIdentity pins that strict handling +// refuses registrations whose remote address carries no node ID. +func TestRegisterStrictRequiresCallerIdentity(t *testing.T) { + t.Parallel() + s := New(&fakeListener{}, "") + s.SetStrictRegister(true) + + addr := protocol.Addr{Network: 0, Node: 0x99} + reqs := []Request{ + {Command: "REGISTER", RecordType: RecordA, Name: "a", Address: addr.String()}, + {Command: "REGISTER", RecordType: RecordN, Name: "n", NetID: 1}, + {Command: "REGISTER", RecordType: RecordS, Name: "s", Address: addr.String(), NetID: 1, Port: 80}, + } + for _, req := range reqs { + if got := s.handleRequest(req, nil); !strings.HasPrefix(got, "ERR ") { + t.Errorf("REGISTER %s with no caller identity: got %q, want ERR", req.RecordType, got) + } + if got := s.handleRequest(req, stringAddr{s: "127.0.0.1:8080"}); !strings.HasPrefix(got, "ERR ") { + t.Errorf("REGISTER %s from non-pilot addr: got %q, want ERR", req.RecordType, got) + } + } + if _, err := s.store.LookupA("a"); err == nil { + t.Error("A record was stored despite refusal") + } + if _, err := s.store.LookupN("n"); err == nil { + t.Error("N record was stored despite refusal") + } + if svcs := s.store.LookupS(1, 80); len(svcs) != 0 { + t.Errorf("S record was stored despite refusal: %+v", svcs) + } +} + +// TestRegisterA_StrictMatchesCallerNode pins that with strict handling on +// the caller node is actually resolved from the "addr:port" remote address +// form, so the existing self-registration check fires. +func TestRegisterA_StrictMatchesCallerNode(t *testing.T) { + t.Parallel() + s := New(&fakeListener{}, "") + s.SetStrictRegister(true) + + foreign := protocol.Addr{Network: 0, Node: 0x99} + req := Request{Command: "REGISTER", RecordType: RecordA, Name: "bob", Address: foreign.String()} + if got := s.handleRequest(req, callerAddr(1)); !strings.HasPrefix(got, "ERR ") { + t.Errorf("REGISTER A for another node: got %q, want ERR", got) + } + + own := protocol.Addr{Network: 0, Node: 1} + req = Request{Command: "REGISTER", RecordType: RecordA, Name: "self", Address: own.String()} + if got := s.handleRequest(req, callerAddr(1)); got != "OK" { + t.Errorf("REGISTER A for own node: got %q", got) + } + if addr, err := s.store.LookupA("self"); err != nil || addr != own { + t.Errorf("LookupA(self): got (%v, %v)", addr, err) + } +} + +func TestStrictRegisterFromEnv(t *testing.T) { + for _, tc := range []struct { + val string + want bool + }{ + {"", false}, + {"0", false}, + {"no", false}, + {"1", true}, + {"true", true}, + {"TRUE", true}, + } { + t.Setenv(StrictRegisterEnv, tc.val) + s := New(&fakeListener{}, "") + if got := s.StrictRegister(); got != tc.want { + t.Errorf("%s=%q: got %v, want %v", StrictRegisterEnv, tc.val, got, tc.want) + } + } +} + +// TestRegisterNOwned_PersistsOwner pins that the owning node survives a +// snapshot round-trip, so enabling strict handling later has the data it +// needs. +func TestRegisterNOwned_PersistsOwner(t *testing.T) { + t.Parallel() + dir := t.TempDir() + path := filepath.Join(dir, "ns.json") + + rs := NewRecordStore() + rs.SetStorePath(path) + if !rs.RegisterNOwned("mynet", 5, 0x1234, false) { + t.Fatal("RegisterNOwned returned false") + } + rs.Close() + + rs2 := NewRecordStore() + defer rs2.Close() + rs2.SetStorePath(path) + if netID, err := rs2.LookupN("mynet"); err != nil || netID != 5 { + t.Fatalf("LookupN after reload: got (%d, %v)", netID, err) + } + if rs2.RegisterNOwned("mynet", 9, 0x5678, true) { + t.Error("reloaded record accepted a different owner under enforcement") + } + if !rs2.RegisterNOwned("mynet", 9, 0x1234, true) { + t.Error("reloaded record refused its own owner") + } +}