diff --git a/encode.go b/encode.go index de3cf83..0c010ea 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,14 @@ 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() + dimsInitialized := h.Len() > 0 h.layers = make([]*layer[K], nLayers) for i := 0; i < nLayers; i++ { @@ -215,6 +229,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 +242,20 @@ 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 !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++ { @@ -252,7 +283,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..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" @@ -180,6 +181,214 @@ 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{ + "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) + }, + "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 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]()