From d8a8038a67c8d94917ddc9421aab6f56d626c841 Mon Sep 17 00:00:00 2001 From: Tyagiquamar Date: Tue, 25 Aug 2026 19:27:53 +0530 Subject: [PATCH 1/2] fix: reject invalid graphs in Import instead of panicking later Import trusted the encoded payload structurally: - a neighbor key absent from its layer was stored as a nil *layerNode, so the first Search over the imported graph panicked with a nil pointer dereference - vectors of mismatched dimensionality were accepted (within one file, or against vectors already in the graph), making Search panic inside the distance function depending on which node became the entry point - negative layer/node/neighbor/vector counts reached make() and panicked during decoding All four are rejected with errors now, matching the existing version and distance-function validation. Valid Export output never triggers any of the new checks: neighbors are always intra-layer and assertDims keeps vector dimensions uniform. --- encode.go | 30 +++++++++- encode_test.go | 148 +++++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 177 insertions(+), 1 deletion(-) diff --git a/encode.go b/encode.go index de3cf83..ee68d90 100644 --- a/encode.go +++ b/encode.go @@ -38,6 +38,9 @@ func binaryRead(r io.Reader, data interface{}) (int, error) { if err != nil { return 0, err } + if ln < 0 { + return 0, fmt.Errorf("invalid string length: %d", ln) + } s := make([]byte, ln) _, err = binaryRead(r, &s) @@ -50,6 +53,9 @@ func binaryRead(r io.Reader, data interface{}) (int, error) { if err != nil { return 0, err } + if ln < 0 { + return 0, fmt.Errorf("invalid vector length: %d", ln) + } *v = make([]float32, ln) return binary.Size(*v), binary.Read(r, byteOrder, *v) @@ -207,6 +213,13 @@ func (h *Graph[K]) Import(r io.Reader) error { if err != nil { return err } + if nLayers < 0 { + return fmt.Errorf("invalid number of layers: %d", nLayers) + } + + // Every vector must share the dimensionality of the vectors already in + // the graph, or of the first vector decoded if it is empty. + dims := h.Dims() h.layers = make([]*layer[K], nLayers) for i := 0; i < nLayers; i++ { @@ -215,6 +228,9 @@ func (h *Graph[K]) Import(r io.Reader) error { if err != nil { return err } + if nNodes < 0 { + return fmt.Errorf("invalid number of nodes in layer %d: %d", i, nNodes) + } nodes := make(map[K]*layerNode[K], nNodes) for j := 0; j < nNodes; j++ { @@ -225,6 +241,14 @@ func (h *Graph[K]) Import(r io.Reader) error { if err != nil { return fmt.Errorf("decoding node %d: %w", j, err) } + if nNeighbors < 0 { + return fmt.Errorf("invalid neighbor count for node %v: %d", key, nNeighbors) + } + if dims == 0 { + dims = len(vec) + } else if len(vec) != dims { + return fmt.Errorf("node %v has a %d-dimensional vector, expected %d dimensions", key, len(vec), dims) + } neighbors := make([]K, nNeighbors) for k := 0; k < nNeighbors; k++ { @@ -252,7 +276,11 @@ func (h *Graph[K]) Import(r io.Reader) error { // Fill in neighbor pointers for _, node := range nodes { for key := range node.neighbors { - node.neighbors[key] = nodes[key] + target, ok := nodes[key] + if !ok { + return fmt.Errorf("node %v has neighbor %v, but it is not present in its layer", node.Key, key) + } + node.neighbors[key] = target } } h.layers[i] = &layer[K]{nodes: nodes} diff --git a/encode_test.go b/encode_test.go index b19210e..913ed3a 100644 --- a/encode_test.go +++ b/encode_test.go @@ -180,6 +180,154 @@ func TestSavedGraph(t *testing.T) { const benchGraphSize = 100 +// encodeTestPayload encodes vals using the wire format of Export. +func encodeTestPayload(t *testing.T, vals ...any) *bytes.Buffer { + t.Helper() + buf := &bytes.Buffer{} + _, err := multiBinaryWrite(buf, vals...) + require.NoError(t, err) + return buf +} + +func TestGraph_ImportRejectsDanglingNeighbor(t *testing.T) { + t.Parallel() + + // One layer holding a single node whose only neighbor references a key + // that is absent from the layer. Importing this must fail rather than + // produce a graph containing nil neighbors that panics on Search. + buf := encodeTestPayload( + t, + encodingVersion, + 16, + 0.25, + 20, + "cosine", + 1, // layers + 1, // nodes + 1, // key + Vector{1, 0, 0}, // value + 1, // neighbors + 999, // dangling neighbor key + ) + + g := NewGraph[int]() + err := g.Import(buf) + require.ErrorContains(t, err, "not present") +} + +func TestGraph_ImportRejectsInconsistentDims(t *testing.T) { + t.Parallel() + + t.Run("WithinFile", func(t *testing.T) { + t.Parallel() + + // Two nodes whose vectors have different dimensions must be + // rejected: searching such a graph panics inside the distance + // function depending on which node becomes the entry point. + buf := encodeTestPayload( + t, + encodingVersion, + 16, + 0.25, + 20, + "cosine", + 1, // layers + 2, // nodes + 1, // key + Vector{1, 0, 0}, // value + 0, // neighbors + 2, // key + Vector{1, 0}, // value + 0, // neighbors + ) + + g := NewGraph[int]() + err := g.Import(buf) + require.ErrorContains(t, err, "dimensions") + }) + + t.Run("AgainstGraph", func(t *testing.T) { + t.Parallel() + + // Like Add, importing into a populated graph requires matching + // dimensionality. + g := NewGraph[int]() + g.Add(MakeNode(1, Vector{1})) + + buf := encodeTestPayload( + t, + encodingVersion, + 16, + 0.25, + 20, + "cosine", + 1, // layers + 1, // nodes + 2, // key + Vector{1, 0}, // value + 0, // neighbors + ) + + err := g.Import(buf) + require.ErrorContains(t, err, "dimensions") + }) +} + +func TestGraph_ImportRejectsInvalidCounts(t *testing.T) { + t.Parallel() + + tests := map[string]func(t *testing.T) *bytes.Buffer{ + "negativeLayers": func(t *testing.T) *bytes.Buffer { + return encodeTestPayload(t, encodingVersion, 16, 0.25, 20, "cosine", -1) + }, + "negativeNodes": func(t *testing.T) *bytes.Buffer { + return encodeTestPayload(t, encodingVersion, 16, 0.25, 20, "cosine", 1, -2) + }, + "negativeNeighbors": func(t *testing.T) *bytes.Buffer { + return encodeTestPayload( + t, + encodingVersion, + 16, + 0.25, + 20, + "cosine", + 1, // layers + 1, // nodes + 1, // key + Vector{1}, // value + -3, // neighbors + ) + }, + "negativeVectorLength": func(t *testing.T) *bytes.Buffer { + return encodeTestPayload( + t, + encodingVersion, + 16, + 0.25, + 20, + "cosine", + 1, // layers + 1, // nodes + 1, // key + -4, // vector length + 0, // neighbors + ) + }, + } + + for name, encode := range tests { + name, encode := name, encode + t.Run(name, func(t *testing.T) { + t.Parallel() + + payload := encode(t) + g := NewGraph[int]() + err := g.Import(payload) + require.Error(t, err) + }) + } +} + func BenchmarkGraph_Import(b *testing.B) { b.ReportAllocs() g := newTestGraph[int]() From 42f48a0e7efc22fc41f15b5a28b32c49c9a338d2 Mon Sep 17 00:00:00 2001 From: Tyagiquamar Date: Thu, 8 Oct 2026 22:37:53 +0530 Subject: [PATCH 2/2] fix: address import validation review feedback (#25) --- encode.go | 11 +++++++-- encode_test.go | 61 ++++++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 70 insertions(+), 2 deletions(-) diff --git a/encode.go b/encode.go index ee68d90..0c010ea 100644 --- a/encode.go +++ b/encode.go @@ -220,6 +220,7 @@ func (h *Graph[K]) Import(r io.Reader) error { // Every vector must share the dimensionality of the vectors already in // the graph, or of the first vector decoded if it is empty. dims := h.Dims() + dimsInitialized := h.Len() > 0 h.layers = make([]*layer[K], nLayers) for i := 0; i < nLayers; i++ { @@ -228,7 +229,7 @@ func (h *Graph[K]) Import(r io.Reader) error { if err != nil { return err } - if nNodes < 0 { + if nNodes <= 0 { return fmt.Errorf("invalid number of nodes in layer %d: %d", i, nNodes) } @@ -244,11 +245,17 @@ func (h *Graph[K]) Import(r io.Reader) error { if nNeighbors < 0 { return fmt.Errorf("invalid neighbor count for node %v: %d", key, nNeighbors) } - if dims == 0 { + if !dimsInitialized { dims = len(vec) + dimsInitialized = true } else if len(vec) != dims { return fmt.Errorf("node %v has a %d-dimensional vector, expected %d dimensions", key, len(vec), dims) } + if i > 0 { + if _, ok := h.layers[i-1].nodes[key]; !ok { + return fmt.Errorf("node %v in layer %d is not present in the preceding layer", key, i) + } + } neighbors := make([]K, nNeighbors) for k := 0; k < nNeighbors; k++ { diff --git a/encode_test.go b/encode_test.go index 913ed3a..c6aa959 100644 --- a/encode_test.go +++ b/encode_test.go @@ -3,6 +3,7 @@ package hnsw import ( "bytes" "cmp" + "fmt" "testing" "github.com/stretchr/testify/require" @@ -277,6 +278,12 @@ func TestGraph_ImportRejectsInvalidCounts(t *testing.T) { t.Parallel() tests := map[string]func(t *testing.T) *bytes.Buffer{ + "negativeStringLength": func(t *testing.T) *bytes.Buffer { + return encodeTestPayload(t, encodingVersion, 16, 0.25, 20, -1) + }, + "emptyLayer": func(t *testing.T) *bytes.Buffer { + return encodeTestPayload(t, encodingVersion, 16, 0.25, 20, "cosine", 1, 0) + }, "negativeLayers": func(t *testing.T) *bytes.Buffer { return encodeTestPayload(t, encodingVersion, 16, 0.25, 20, "cosine", -1) }, @@ -328,6 +335,60 @@ func TestGraph_ImportRejectsInvalidCounts(t *testing.T) { } } +func TestGraph_ImportRejectsMissingLowerLayerNode(t *testing.T) { + t.Parallel() + for _, upperKey := range []int{1, 2} { + t.Run(fmt.Sprint(upperKey), func(t *testing.T) { + buf := encodeTestPayload(t, encodingVersion, 16, 0.25, 20, "cosine", + 3, 1, 1, Vector{1}, 0, 1, 1, Vector{1}, 0, 1, upperKey, Vector{1}, 0) + g := NewGraph[int]() + err := g.Import(buf) + if upperKey == 1 { + require.NoError(t, err) + require.Equal(t, 1, g.Search(Vector{1}, 1)[0].Key) + } else { + require.ErrorContains(t, err, "preceding layer") + } + }) + } +} + +func TestGraph_ImportZeroDims(t *testing.T) { + t.Parallel() + t.Run("AgainstGraph", func(t *testing.T) { + g := NewGraph[int]() + g.Add(MakeNode(3, Vector{})) + buf := encodeTestPayload(t, encodingVersion, 16, 0.25, 20, "cosine", + 1, 1, 1, Vector{1}, 0) + require.ErrorContains(t, g.Import(buf), "dimensions") + }) + t.Run("EmptyGraph", func(t *testing.T) { + g := NewGraph[int]() + buf := encodeTestPayload(t, encodingVersion, 16, 0.25, 20, "cosine", 0) + require.NoError(t, g.Import(buf)) + require.Zero(t, g.Dims()) + require.Empty(t, g.Search(Vector{}, 1)) + }) + for _, populated := range []bool{false, true} { + for _, dims := range []int{0, 1} { + t.Run(fmt.Sprintf("populated=%t/dims=%d", populated, dims), func(t *testing.T) { + g := NewGraph[int]() + if populated { + g.Add(MakeNode(3, Vector{})) + } + buf := encodeTestPayload(t, encodingVersion, 16, 0.25, 20, "cosine", + 1, 2, 1, Vector{}, 0, 2, make(Vector, dims), 0) + err := g.Import(buf) + if dims == 0 { + require.NoError(t, err) + } else { + require.ErrorContains(t, err, "dimensions") + } + }) + } + } +} + func BenchmarkGraph_Import(b *testing.B) { b.ReportAllocs() g := newTestGraph[int]()