Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 29 additions & 1 deletion encode.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}
Comment on lines +41 to +43

s := make([]byte, ln)
_, err = binaryRead(r, &s)
Expand All @@ -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)
Expand Down Expand Up @@ -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++ {
Expand All @@ -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)
}
Comment on lines +231 to +233

nodes := make(map[K]*layerNode[K], nNodes)
for j := 0; j < nNodes; j++ {
Expand All @@ -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)
Comment on lines +247 to +250
}

neighbors := make([]K, nNeighbors)
for k := 0; k < nNeighbors; k++ {
Expand Down Expand Up @@ -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
Comment on lines +279 to +283
}
}
h.layers[i] = &layer[K]{nodes: nodes}
Expand Down
148 changes: 148 additions & 0 deletions encode_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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]()
Expand Down
Loading