diff --git a/go.mod b/go.mod index 3c3da45..25485ed 100644 --- a/go.mod +++ b/go.mod @@ -8,13 +8,13 @@ require github.com/creack/pty v1.1.24 require ( github.com/gorilla/websocket v1.5.3 - google.golang.org/grpc v1.72.2 - google.golang.org/protobuf v1.36.6 + google.golang.org/grpc v1.82.1 + google.golang.org/protobuf v1.36.11 ) require ( - golang.org/x/net v0.35.0 // indirect - golang.org/x/sys v0.30.0 // indirect - golang.org/x/text v0.22.0 // indirect - google.golang.org/genproto/googleapis/rpc v0.0.0-20250218202821-56aae31c358a // indirect + golang.org/x/net v0.53.0 // indirect + golang.org/x/sys v0.43.0 // indirect + golang.org/x/text v0.36.0 // indirect + google.golang.org/genproto/googleapis/rpc v0.0.0-20260414002931-afd174a4e478 // indirect ) diff --git a/go.sum b/go.sum index 3a8158a..eab46a9 100644 --- a/go.sum +++ b/go.sum @@ -1,40 +1,44 @@ github.com/LatticeNet/lattice-sdk v0.2.18-0.20260722123932-4a318f246d23 h1:2qpbnG8jO9lOlVRpp4OCeami34HUyV7VGmp3vj4IRQI= github.com/LatticeNet/lattice-sdk v0.2.18-0.20260722123932-4a318f246d23/go.mod h1:7ENUQ4EoS/TSW/eNomCGfZGliUPJZ46uAvp7dVcEXoE= +github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= +github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= github.com/creack/pty v1.1.24 h1:bJrF4RRfyJnbTJqzRLHzcGaZK1NeM5kTC9jGgovnR1s= github.com/creack/pty v1.1.24/go.mod h1:08sCNb52WyoAwi2QDyzUCTgcvVFhUzewun7wtTfvcwE= -github.com/go-logr/logr v1.4.2 h1:6pFjapn8bFcIbiKo3XT4j/BhANplGihG6tvd+8rYgrY= -github.com/go-logr/logr v1.4.2/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY= +github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI= +github.com/go-logr/logr v1.4.3/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY= github.com/go-logr/stdr v1.2.2 h1:hSWxHoqTgW2S2qGc0LTAI563KZ5YKYRhT3MFKZMbjag= github.com/go-logr/stdr v1.2.2/go.mod h1:mMo/vtBO5dYbehREoey6XUKy/eSumjCCveDpRre4VKE= github.com/golang/protobuf v1.5.4 h1:i7eJL8qZTpSEXOPTxNKhASYpMn+8e5Q6AdndVa1dWek= github.com/golang/protobuf v1.5.4/go.mod h1:lnTiLA8Wa4RWRcIUkrtSVa5nRhsEGBg48fD6rSs7xps= -github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI= -github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY= +github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= +github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg= github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE= -go.opentelemetry.io/auto/sdk v1.1.0 h1:cH53jehLUN6UFLY71z+NDOiNJqDdPRaXzTel0sJySYA= -go.opentelemetry.io/auto/sdk v1.1.0/go.mod h1:3wSPjt5PWp2RhlCcmmOial7AvC4DQqZb7a7wCow3W8A= -go.opentelemetry.io/otel v1.34.0 h1:zRLXxLCgL1WyKsPVrgbSdMN4c0FMkDAskSTQP+0hdUY= -go.opentelemetry.io/otel v1.34.0/go.mod h1:OWFPOQ+h4G8xpyjgqo4SxJYdDQ/qmRH+wivy7zzx9oI= -go.opentelemetry.io/otel/metric v1.34.0 h1:+eTR3U0MyfWjRDhmFMxe2SsW64QrZ84AOhvqS7Y+PoQ= -go.opentelemetry.io/otel/metric v1.34.0/go.mod h1:CEDrp0fy2D0MvkXE+dPV7cMi8tWZwX3dmaIhwPOaqHE= -go.opentelemetry.io/otel/sdk v1.34.0 h1:95zS4k/2GOy069d321O8jWgYsW3MzVV+KuSPKp7Wr1A= -go.opentelemetry.io/otel/sdk v1.34.0/go.mod h1:0e/pNiaMAqaykJGKbi+tSjWfNNHMTxoC9qANsCzbyxU= -go.opentelemetry.io/otel/sdk/metric v1.34.0 h1:5CeK9ujjbFVL5c1PhLuStg1wxA7vQv7ce1EK0Gyvahk= -go.opentelemetry.io/otel/sdk/metric v1.34.0/go.mod h1:jQ/r8Ze28zRKoNRdkjCZxfs6YvBTG1+YIqyFVFYec5w= -go.opentelemetry.io/otel/trace v1.34.0 h1:+ouXS2V8Rd4hp4580a8q23bg0azF2nI8cqLYnC8mh/k= -go.opentelemetry.io/otel/trace v1.34.0/go.mod h1:Svm7lSjQD7kG7KJ/MUHPVXSDGz2OX4h0M2jHBhmSfRE= -golang.org/x/net v0.35.0 h1:T5GQRQb2y08kTAByq9L4/bz8cipCdA8FbRTXewonqY8= -golang.org/x/net v0.35.0/go.mod h1:EglIi67kWsHKlRzzVMUD93VMSWGFOMSZgxFjparz1Qk= -golang.org/x/sys v0.30.0 h1:QjkSwP/36a20jFYWkSue1YwXzLmsV5Gfq7Eiy72C1uc= -golang.org/x/sys v0.30.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA= -golang.org/x/text v0.22.0 h1:bofq7m3/HAFvbF51jz3Q9wLg3jkvSPuiZu/pD1XwgtM= -golang.org/x/text v0.22.0/go.mod h1:YRoo4H8PVmsu+E3Ou7cqLVH8oXWIHVoX0jqUWALQhfY= -google.golang.org/genproto/googleapis/rpc v0.0.0-20250218202821-56aae31c358a h1:51aaUVRocpvUOSQKM6Q7VuoaktNIaMCLuhZB6DKksq4= -google.golang.org/genproto/googleapis/rpc v0.0.0-20250218202821-56aae31c358a/go.mod h1:uRxBH1mhmO8PGhU89cMcHaXKZqO+OfakD8QQO0oYwlQ= -google.golang.org/grpc v1.72.2 h1:TdbGzwb82ty4OusHWepvFWGLgIbNo1/SUynEN0ssqv8= -google.golang.org/grpc v1.72.2/go.mod h1:wH5Aktxcg25y1I3w7H69nHfXdOG3UiadoBtjh3izSDM= -google.golang.org/protobuf v1.36.6 h1:z1NpPI8ku2WgiWnf+t9wTPsn6eP1L7ksHUlkfLvd9xY= -google.golang.org/protobuf v1.36.6/go.mod h1:jduwjTPXsFjZGTmRluh+L6NjiWu7pchiJ2/5YcXBHnY= +go.opentelemetry.io/auto/sdk v1.2.1 h1:jXsnJ4Lmnqd11kwkBV2LgLoFMZKizbCi5fNZ/ipaZ64= +go.opentelemetry.io/auto/sdk v1.2.1/go.mod h1:KRTj+aOaElaLi+wW1kO/DZRXwkF4C5xPbEe3ZiIhN7Y= +go.opentelemetry.io/otel v1.43.0 h1:mYIM03dnh5zfN7HautFE4ieIig9amkNANT+xcVxAj9I= +go.opentelemetry.io/otel v1.43.0/go.mod h1:JuG+u74mvjvcm8vj8pI5XiHy1zDeoCS2LB1spIq7Ay0= +go.opentelemetry.io/otel/metric v1.43.0 h1:d7638QeInOnuwOONPp4JAOGfbCEpYb+K6DVWvdxGzgM= +go.opentelemetry.io/otel/metric v1.43.0/go.mod h1:RDnPtIxvqlgO8GRW18W6Z/4P462ldprJtfxHxyKd2PY= +go.opentelemetry.io/otel/sdk v1.43.0 h1:pi5mE86i5rTeLXqoF/hhiBtUNcrAGHLKQdhg4h4V9Dg= +go.opentelemetry.io/otel/sdk v1.43.0/go.mod h1:P+IkVU3iWukmiit/Yf9AWvpyRDlUeBaRg6Y+C58QHzg= +go.opentelemetry.io/otel/sdk/metric v1.43.0 h1:S88dyqXjJkuBNLeMcVPRFXpRw2fuwdvfCGLEo89fDkw= +go.opentelemetry.io/otel/sdk/metric v1.43.0/go.mod h1:C/RJtwSEJ5hzTiUz5pXF1kILHStzb9zFlIEe85bhj6A= +go.opentelemetry.io/otel/trace v1.43.0 h1:BkNrHpup+4k4w+ZZ86CZoHHEkohws8AY+WTX09nk+3A= +go.opentelemetry.io/otel/trace v1.43.0/go.mod h1:/QJhyVBUUswCphDVxq+8mld+AvhXZLhe+8WVFxiFff0= +golang.org/x/net v0.53.0 h1:d+qAbo5L0orcWAr0a9JweQpjXF19LMXJE8Ey7hwOdUA= +golang.org/x/net v0.53.0/go.mod h1:JvMuJH7rrdiCfbeHoo3fCQU24Lf5JJwT9W3sJFulfgs= +golang.org/x/sys v0.43.0 h1:Rlag2XtaFTxp19wS8MXlJwTvoh8ArU6ezoyFsMyCTNI= +golang.org/x/sys v0.43.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/text v0.36.0 h1:JfKh3XmcRPqZPKevfXVpI1wXPTqbkE5f7JA92a55Yxg= +golang.org/x/text v0.36.0/go.mod h1:NIdBknypM8iqVmPiuco0Dh6P5Jcdk8lJL0CUebqK164= +gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4= +gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E= +google.golang.org/genproto/googleapis/rpc v0.0.0-20260414002931-afd174a4e478 h1:RmoJA1ujG+/lRGNfUnOMfhCy5EipVMyvUE+KNbPbTlw= +google.golang.org/genproto/googleapis/rpc v0.0.0-20260414002931-afd174a4e478/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8= +google.golang.org/grpc v1.82.1 h1:NnAxzGRA0677vCa4BUkOAnO5+FfQqVl9iUXeD0IqcGE= +google.golang.org/grpc v1.82.1/go.mod h1:yzTZ1TB1Z3SG+LIYaI+WiE8D5+PZ3ArnrSp8zF3+/ZA= +google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE= +google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= diff --git a/internal/guardreality/collect.go b/internal/guardreality/collect.go new file mode 100644 index 0000000..6afabee --- /dev/null +++ b/internal/guardreality/collect.go @@ -0,0 +1,477 @@ +// Package guardreality collects the read-only node facts NetGuard uses for +// reality-first firewall authoring. The package only observes host state; it +// never mutates nftables, interfaces, or processes. +package guardreality + +import ( + "bytes" + "context" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "fmt" + "io" + "net" + "os/exec" + "regexp" + "sort" + "strconv" + "strings" + "time" + + "github.com/LatticeNet/lattice-sdk/model" +) + +const ( + defaultTimeout = 5 * time.Second + maxOutputBytes = 1 << 20 + managedFamily = "inet" + managedTable = "lattice_guard" +) + +// Runner executes one command without a shell. Tests inject it so parser +// coverage never depends on the local host's nftables, ss, or ip state. +type Runner func(ctx context.Context, name string, args ...string) ([]byte, error) + +// Source configures guard reality collection. +type Source struct { + SSBinary string + IPBinary string + NFTBinary string + Timeout time.Duration + Now func() time.Time + Runner Runner +} + +// Collect runs the read-only guard reality commands and normalizes their output +// into the shared SDK model. The caller-supplied node id wins over anything a +// command could report. +func Collect(ctx context.Context, source Source, nodeID string) (model.GuardNodeReality, error) { + nodeID = strings.TrimSpace(nodeID) + if nodeID == "" { + return model.GuardNodeReality{}, fmt.Errorf("node id is required") + } + if ctx == nil { + ctx = context.Background() + } + timeout := source.Timeout + if timeout <= 0 { + timeout = defaultTimeout + } + run := source.Runner + if run == nil { + run = runBoundedCommand + } + ssBinary := firstNonEmpty(source.SSBinary, "ss") + ipBinary := firstNonEmpty(source.IPBinary, "ip") + nftBinary := firstNonEmpty(source.NFTBinary, "nft") + + ssOut, err := runStep(ctx, timeout, run, ssBinary, "-tulpnH") + if err != nil { + return model.GuardNodeReality{}, err + } + listeners, err := ParseSSListeners(ssOut) + if err != nil { + return model.GuardNodeReality{}, fmt.Errorf("parse ss listeners: %w", err) + } + + ipOut, err := runStep(ctx, timeout, run, ipBinary, "-j", "addr") + if err != nil { + return model.GuardNodeReality{}, err + } + interfaces, err := ParseIPAddr(ipOut) + if err != nil { + return model.GuardNodeReality{}, fmt.Errorf("parse ip addr: %w", err) + } + + rulesetOut, err := runStep(ctx, timeout, run, nftBinary, "-j", "list", "ruleset") + if err != nil { + return model.GuardNodeReality{}, err + } + managedSHA, foreignTables, err := ParseNFTRuleset(rulesetOut) + if err != nil { + return model.GuardNodeReality{}, fmt.Errorf("parse nft ruleset: %w", err) + } + + versionOut, err := runStep(ctx, timeout, run, nftBinary, "--version") + if err != nil { + return model.GuardNodeReality{}, err + } + at := time.Now().UTC() + if source.Now != nil { + at = source.Now().UTC() + } + return model.GuardNodeReality{ + NodeID: nodeID, + Listeners: listeners, + Interfaces: interfaces, + ManagedSHA: managedSHA, + ForeignTables: foreignTables, + NFTVersion: ParseNFTVersion(versionOut), + CollectedAt: at, + }, nil +} + +func runStep(ctx context.Context, timeout time.Duration, run Runner, name string, args ...string) ([]byte, error) { + stepCtx, cancel := context.WithTimeout(ctx, timeout) + defer cancel() + out, err := run(stepCtx, name, args...) + if err != nil { + return nil, fmt.Errorf("%s %s: %w", name, strings.Join(args, " "), err) + } + return out, nil +} + +var ssProcessRe = regexp.MustCompile(`\(\("([^"]+)"`) + +// ParseSSListeners parses `ss -tulpnH` output into deterministic listener facts. +func ParseSSListeners(raw []byte) ([]model.GuardListener, error) { + lines := strings.Split(string(raw), "\n") + seen := map[string]struct{}{} + out := make([]model.GuardListener, 0, len(lines)) + for _, line := range lines { + line = strings.TrimSpace(line) + if line == "" { + continue + } + fields := strings.Fields(line) + if len(fields) < 5 { + continue + } + proto := strings.ToLower(fields[0]) + if proto != "tcp" && proto != "udp" { + continue + } + addr, port, ok := parseHostPort(fields[4]) + if !ok || port <= 0 || port > 65535 { + continue + } + process := "" + if match := ssProcessRe.FindStringSubmatch(line); len(match) == 2 { + process = trimBounded(match[1], 128) + } + listener := model.GuardListener{ + Protocol: proto, + Port: port, + Address: trimBounded(addr, 256), + Process: process, + } + key := fmt.Sprintf("%s/%d/%s/%s", listener.Protocol, listener.Port, listener.Address, listener.Process) + if _, ok := seen[key]; ok { + continue + } + seen[key] = struct{}{} + out = append(out, listener) + } + sort.Slice(out, func(i, j int) bool { + if out[i].Protocol != out[j].Protocol { + return out[i].Protocol < out[j].Protocol + } + if out[i].Port != out[j].Port { + return out[i].Port < out[j].Port + } + if out[i].Address != out[j].Address { + return out[i].Address < out[j].Address + } + return out[i].Process < out[j].Process + }) + return out, nil +} + +func parseHostPort(value string) (string, int, bool) { + value = strings.TrimSpace(value) + if value == "" { + return "", 0, false + } + var host, portText string + if strings.HasPrefix(value, "[") { + end := strings.LastIndex(value, "]:") + if end < 0 { + return "", 0, false + } + host = value[1:end] + portText = value[end+2:] + } else { + idx := strings.LastIndex(value, ":") + if idx < 0 { + return "", 0, false + } + host = value[:idx] + portText = value[idx+1:] + } + if zone := strings.LastIndex(host, "%"); zone >= 0 { + host = host[:zone] + } + port, err := strconv.Atoi(strings.TrimSpace(portText)) + if err != nil { + return "", 0, false + } + return strings.TrimSpace(host), port, true +} + +type ipAddrEntry struct { + IfName string `json:"ifname"` + Flags []string `json:"flags"` + OperState string `json:"operstate"` + AddrInfo []ipAddrInfo `json:"addr_info"` +} + +type ipAddrInfo struct { + Local string `json:"local"` + PrefixLen *int `json:"prefixlen"` +} + +// ParseIPAddr parses `ip -j addr` output. +func ParseIPAddr(raw []byte) ([]model.GuardInterface, error) { + var entries []ipAddrEntry + if err := json.Unmarshal(bytes.TrimSpace(raw), &entries); err != nil { + return nil, err + } + out := make([]model.GuardInterface, 0, len(entries)) + for _, entry := range entries { + name := strings.TrimSpace(entry.IfName) + if name == "" { + continue + } + addresses := make([]string, 0, len(entry.AddrInfo)) + for _, addr := range entry.AddrInfo { + local := strings.TrimSpace(addr.Local) + ip := net.ParseIP(local) + if ip == nil { + continue + } + canonical := ip.String() + maxPrefix := 128 + if ip.To4() != nil { + maxPrefix = 32 + } + if addr.PrefixLen != nil && *addr.PrefixLen >= 0 && *addr.PrefixLen <= maxPrefix { + canonical = fmt.Sprintf("%s/%d", canonical, *addr.PrefixLen) + } + addresses = append(addresses, canonical) + } + sort.Strings(addresses) + out = append(out, model.GuardInterface{ + Name: trimBounded(name, 128), + Addresses: uniqueStrings(addresses), + Up: ifaceIsUp(entry), + }) + } + sort.Slice(out, func(i, j int) bool { return out[i].Name < out[j].Name }) + return out, nil +} + +func ifaceIsUp(entry ipAddrEntry) bool { + for _, flag := range entry.Flags { + if strings.EqualFold(flag, "UP") { + return true + } + } + return strings.EqualFold(entry.OperState, "UP") +} + +// ParseNFTRuleset parses `nft -j list ruleset`, returning a deterministic hash +// of the managed lattice_guard table and sorted summaries of foreign tables. +func ParseNFTRuleset(raw []byte) (string, []string, error) { + var payload struct { + NFTables []map[string]any `json:"nftables"` + } + if err := json.Unmarshal(bytes.TrimSpace(raw), &payload); err != nil { + return "", nil, err + } + foreign := map[string]struct{}{} + managedObjects := make([]map[string]any, 0) + for _, entry := range payload.NFTables { + for kind, rawBody := range entry { + body, ok := rawBody.(map[string]any) + if !ok { + continue + } + if kind == "table" { + family := stringField(body, "family") + name := stringField(body, "name") + if family == "" || name == "" { + continue + } + if family != managedFamily || name != managedTable { + foreign[family+" "+name] = struct{}{} + continue + } + } + if nftObjectBelongsToManaged(kind, body) { + managedObjects = append(managedObjects, map[string]any{ + kind: stripNFTVolatile(body), + }) + } + } + } + foreignTables := make([]string, 0, len(foreign)) + for table := range foreign { + foreignTables = append(foreignTables, table) + } + sort.Strings(foreignTables) + if len(managedObjects) == 0 { + return "", foreignTables, nil + } + encoded, err := json.Marshal(managedObjects) + if err != nil { + return "", nil, err + } + sum := sha256.Sum256(encoded) + return hex.EncodeToString(sum[:]), foreignTables, nil +} + +func nftObjectBelongsToManaged(kind string, body map[string]any) bool { + family := stringField(body, "family") + switch kind { + case "table": + return family == managedFamily && stringField(body, "name") == managedTable + default: + return family == managedFamily && stringField(body, "table") == managedTable + } +} + +func stripNFTVolatile(v any) any { + switch typed := v.(type) { + case map[string]any: + out := make(map[string]any, len(typed)) + keys := make([]string, 0, len(typed)) + for key := range typed { + if key == "handle" { + continue + } + keys = append(keys, key) + } + sort.Strings(keys) + for _, key := range keys { + out[key] = stripNFTVolatile(typed[key]) + } + return out + case []any: + out := make([]any, 0, len(typed)) + for _, item := range typed { + out = append(out, stripNFTVolatile(item)) + } + return out + default: + return typed + } +} + +func stringField(body map[string]any, key string) string { + value, _ := body[key].(string) + return strings.TrimSpace(value) +} + +// ParseNFTVersion normalizes `nft --version` output for display. +func ParseNFTVersion(raw []byte) string { + line := "" + for _, candidate := range strings.Split(string(raw), "\n") { + candidate = strings.TrimSpace(candidate) + if candidate != "" { + line = candidate + break + } + } + return trimBounded(line, 128) +} + +func runBoundedCommand(ctx context.Context, name string, args ...string) ([]byte, error) { + cmd := exec.CommandContext(ctx, name, args...) + var stdout limitedBuffer + var stderr limitedBuffer + cmd.Stdout = &stdout + cmd.Stderr = &stderr + if err := cmd.Run(); err != nil { + msg := strings.TrimSpace(stderr.String()) + if msg != "" { + return nil, fmt.Errorf("%w: %s", err, trimBounded(msg, 512)) + } + return nil, err + } + if stdout.Truncated() { + return nil, fmt.Errorf("%s output exceeded %d bytes", name, maxOutputBytes) + } + if stderr.Truncated() { + return nil, fmt.Errorf("%s stderr exceeded %d bytes", name, maxOutputBytes) + } + return stdout.Bytes(), nil +} + +type limitedBuffer struct { + buf bytes.Buffer + truncated bool +} + +func (b *limitedBuffer) Write(p []byte) (int, error) { + if b.buf.Len() < maxOutputBytes { + remaining := maxOutputBytes - b.buf.Len() + if len(p) > remaining { + _, _ = b.buf.Write(p[:remaining]) + b.truncated = true + } else { + _, _ = b.buf.Write(p) + } + } else if len(p) > 0 { + b.truncated = true + } + return len(p), nil +} + +func (b *limitedBuffer) Bytes() []byte { + return b.buf.Bytes() +} + +func (b *limitedBuffer) String() string { + return b.buf.String() +} + +func (b *limitedBuffer) Truncated() bool { + return b.truncated +} + +var _ io.Writer = (*limitedBuffer)(nil) + +func firstNonEmpty(values ...string) string { + for _, value := range values { + if trimmed := strings.TrimSpace(value); trimmed != "" { + return trimmed + } + } + return "" +} + +func uniqueStrings(values []string) []string { + if len(values) == 0 { + return nil + } + seen := map[string]struct{}{} + out := make([]string, 0, len(values)) + for _, value := range values { + value = strings.TrimSpace(value) + if value == "" { + continue + } + if _, ok := seen[value]; ok { + continue + } + seen[value] = struct{}{} + out = append(out, value) + } + return out +} + +func trimBounded(value string, maxRunes int) string { + value = strings.TrimSpace(value) + value = strings.Map(func(r rune) rune { + if r < 32 && r != '\t' { + return -1 + } + return r + }, value) + runes := []rune(value) + if len(runes) > maxRunes { + return string(runes[:maxRunes]) + } + return value +} diff --git a/internal/guardreality/collect_test.go b/internal/guardreality/collect_test.go new file mode 100644 index 0000000..81fe903 --- /dev/null +++ b/internal/guardreality/collect_test.go @@ -0,0 +1,196 @@ +package guardreality + +import ( + "context" + "errors" + "reflect" + "strings" + "testing" + "time" + + "github.com/LatticeNet/lattice-sdk/model" +) + +func TestCollectBuildsRealityFromInjectedCommands(t *testing.T) { + fixed := time.Date(2026, 7, 31, 12, 30, 0, 0, time.UTC) + calls := []string{} + runner := func(ctx context.Context, name string, args ...string) ([]byte, error) { + call := name + " " + strings.Join(args, " ") + calls = append(calls, call) + switch call { + case "ss -tulpnH": + return []byte(ssFixture), nil + case "ip -j addr": + return []byte(ipFixture), nil + case "nft -j list ruleset": + return []byte(nftFixture("11", "22")), nil + case "nft --version": + return []byte("nftables v1.0.9 (Old Doc Yak)\n"), nil + default: + t.Fatalf("unexpected live command shape: %s", call) + return nil, nil + } + } + + got, err := Collect(context.Background(), Source{Runner: runner, Now: func() time.Time { return fixed }}, " node-a ") + if err != nil { + t.Fatal(err) + } + if got.NodeID != "node-a" || !got.CollectedAt.Equal(fixed) { + t.Fatalf("identity/time not normalized: %+v", got) + } + wantCalls := []string{"ss -tulpnH", "ip -j addr", "nft -j list ruleset", "nft --version"} + if !reflect.DeepEqual(calls, wantCalls) { + t.Fatalf("unexpected commands:\n got %#v\nwant %#v", calls, wantCalls) + } + wantListeners := []model.GuardListener{ + {Protocol: "tcp", Port: 22, Address: "0.0.0.0", Process: "sshd"}, + {Protocol: "tcp", Port: 443, Address: "::", Process: "nginx"}, + {Protocol: "udp", Port: 41641, Address: "0.0.0.0", Process: "tailscaled"}, + } + if !reflect.DeepEqual(got.Listeners, wantListeners) { + t.Fatalf("listeners:\n got %#v\nwant %#v", got.Listeners, wantListeners) + } + wantIfaces := []model.GuardInterface{ + {Name: "eth0", Addresses: []string{"192.0.2.10/24", "2001:db8::10/64"}, Up: true}, + {Name: "tailscale0", Addresses: []string{"100.64.0.2/32"}, Up: true}, + } + if !reflect.DeepEqual(got.Interfaces, wantIfaces) { + t.Fatalf("interfaces:\n got %#v\nwant %#v", got.Interfaces, wantIfaces) + } + if got.ManagedSHA == "" { + t.Fatal("managed table hash must be set") + } + if !reflect.DeepEqual(got.ForeignTables, []string{"inet ts-input", "ip filter"}) { + t.Fatalf("foreign tables = %#v", got.ForeignTables) + } + if got.NFTVersion != "nftables v1.0.9 (Old Doc Yak)" { + t.Fatalf("nft version = %q", got.NFTVersion) + } +} + +func TestCollectPropagatesCommandFailure(t *testing.T) { + boom := errors.New("missing ip") + _, err := Collect(context.Background(), Source{Runner: func(ctx context.Context, name string, args ...string) ([]byte, error) { + if name == "ss" { + return []byte(ssFixture), nil + } + return nil, boom + }}, "node-a") + if !errors.Is(err, boom) || !strings.Contains(err.Error(), "ip -j addr") { + t.Fatalf("expected ip failure with command context, got %v", err) + } +} + +func TestParseSSListenersSkipsNonNumericPortsAndSorts(t *testing.T) { + got, err := ParseSSListeners([]byte(strings.Join([]string{ + `udp UNCONN 0 0 *:51820 *:* users:(("wg",pid=9,fd=3))`, + `tcp LISTEN 0 4096 127.0.0.1:http 0.0.0.0:* users:(("web",pid=1,fd=3))`, + `tcp LISTEN 0 4096 [fe80::1%eth0]:22 [::]:* users:(("sshd",pid=2,fd=3))`, + `tcp LISTEN 0 4096 [fe80::1%eth0]:22 [::]:* users:(("sshd",pid=2,fd=3))`, + }, "\n"))) + if err != nil { + t.Fatal(err) + } + want := []model.GuardListener{ + {Protocol: "tcp", Port: 22, Address: "fe80::1", Process: "sshd"}, + {Protocol: "udp", Port: 51820, Address: "*", Process: "wg"}, + } + if !reflect.DeepEqual(got, want) { + t.Fatalf("listeners:\n got %#v\nwant %#v", got, want) + } +} + +func TestParseIPAddrNormalizesInterfaces(t *testing.T) { + got, err := ParseIPAddr([]byte(`[ + {"ifname":"tailscale0","flags":["POINTOPOINT","UP"],"addr_info":[{"local":"100.64.0.2","prefixlen":32}]}, + {"ifname":"down0","operstate":"DOWN","addr_info":[{"local":"not-an-ip","prefixlen":24}]}, + {"ifname":"eth0","operstate":"UP","addr_info":[{"local":"2001:db8::10","prefixlen":64},{"local":"192.0.2.10","prefixlen":24},{"local":"192.0.2.10","prefixlen":24},{"local":"198.51.100.9","prefixlen":128}]} + ]`)) + if err != nil { + t.Fatal(err) + } + want := []model.GuardInterface{ + {Name: "down0"}, + {Name: "eth0", Addresses: []string{"192.0.2.10/24", "198.51.100.9", "2001:db8::10/64"}, Up: true}, + {Name: "tailscale0", Addresses: []string{"100.64.0.2/32"}, Up: true}, + } + if !reflect.DeepEqual(got, want) { + t.Fatalf("interfaces:\n got %#v\nwant %#v", got, want) + } +} + +func TestParseNFTRulesetIgnoresHandlesWhenHashingManagedTable(t *testing.T) { + a, foreignA, err := ParseNFTRuleset([]byte(nftFixture("11", "22"))) + if err != nil { + t.Fatal(err) + } + b, foreignB, err := ParseNFTRuleset([]byte(nftFixture("99", "100"))) + if err != nil { + t.Fatal(err) + } + if a == "" || a != b { + t.Fatalf("managed hash should be stable across handle churn: %q vs %q", a, b) + } + wantForeign := []string{"inet ts-input", "ip filter"} + if !reflect.DeepEqual(foreignA, wantForeign) || !reflect.DeepEqual(foreignB, wantForeign) { + t.Fatalf("foreign tables = %#v / %#v", foreignA, foreignB) + } + changed, _, err := ParseNFTRuleset([]byte(strings.ReplaceAll(nftFixture("11", "22"), `"right": 22`, `"right": 2222`))) + if err != nil { + t.Fatal(err) + } + if changed == a { + t.Fatal("managed hash must change when a managed rule changes") + } +} + +func TestParseNFTRulesetWithoutManagedTableReturnsOnlyForeign(t *testing.T) { + hash, foreign, err := ParseNFTRuleset([]byte(`{"nftables":[{"table":{"family":"ip","name":"filter"}}]}`)) + if err != nil { + t.Fatal(err) + } + if hash != "" { + t.Fatalf("managed hash = %q, want empty", hash) + } + if !reflect.DeepEqual(foreign, []string{"ip filter"}) { + t.Fatalf("foreign tables = %#v", foreign) + } +} + +func TestLimitedBufferReportsTruncation(t *testing.T) { + var buf limitedBuffer + chunk := strings.Repeat("x", maxOutputBytes+1) + n, err := buf.Write([]byte(chunk)) + if err != nil || n != len(chunk) { + t.Fatalf("write = %d, %v; want full write count and nil error", n, err) + } + if !buf.Truncated() || len(buf.Bytes()) != maxOutputBytes { + t.Fatalf("truncation not reported correctly: truncated=%v len=%d", buf.Truncated(), len(buf.Bytes())) + } +} + +const ssFixture = ` +tcp LISTEN 0 4096 0.0.0.0:22 0.0.0.0:* users:(("sshd",pid=123,fd=3)) +udp UNCONN 0 0 0.0.0.0:41641 0.0.0.0:* users:(("tailscaled",pid=234,fd=13)) +tcp LISTEN 0 511 [::]:443 [::]:* users:(("nginx",pid=1,fd=6)) +tcp LISTEN 0 4096 127.0.0.1:http 0.0.0.0:* users:(("named-port",pid=2,fd=3)) +` + +const ipFixture = `[ + {"ifname":"eth0","flags":["BROADCAST","MULTICAST","UP"],"addr_info":[{"local":"192.0.2.10","prefixlen":24},{"local":"2001:db8::10","prefixlen":64}]}, + {"ifname":"tailscale0","flags":["POINTOPOINT","UP"],"addr_info":[{"local":"100.64.0.2","prefixlen":32}]} +]` + +func nftFixture(tableHandle, ruleHandle string) string { + return `{ + "nftables": [ + {"metainfo": {"json_schema_version": 1}}, + {"table": {"family": "inet", "name": "lattice_guard", "handle": ` + tableHandle + `}}, + {"chain": {"family": "inet", "table": "lattice_guard", "name": "input", "type": "filter", "hook": "input", "prio": 0, "policy": "drop", "handle": 20}}, + {"rule": {"family": "inet", "table": "lattice_guard", "chain": "input", "expr": [{"match": {"left": {"payload": {"protocol": "tcp", "field": "dport"}}, "op": "==", "right": 22}}, {"accept": null}], "handle": ` + ruleHandle + `}}, + {"table": {"family": "ip", "name": "filter", "handle": 3}}, + {"table": {"family": "inet", "name": "ts-input", "handle": 4}} + ] +}` +}