Skip to content

Commit afd8265

Browse files
committed
perf: reduce memory usage in command suggestions
Use a single-row Levenshtein distance implementation instead of allocating the full DP matrix for every command comparison. The edit distance calculation is unchanged, but temporary allocations are reduced significantly during suggestion generation. Signed-off-by: 0xff-dev <stevenshuang521@gmail.com>
1 parent adbc881 commit afd8265

2 files changed

Lines changed: 251 additions & 17 deletions

File tree

‎cobra.go‎

Lines changed: 30 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -189,37 +189,50 @@ func tmpl(text string) *tmplFunc {
189189
}
190190

191191
// ld compares two strings and returns the levenshtein distance between them.
192+
//
193+
// Only a single row of the distance matrix is kept: computing row i needs
194+
// nothing but row i-1, so the full matrix never has to be materialized. This
195+
// keeps the allocation at O(min(len(s), len(t))) instead of O(len(s)*len(t)).
192196
func ld(s, t string, ignoreCase bool) int {
193197
if ignoreCase {
194198
s = strings.ToLower(s)
195199
t = strings.ToLower(t)
196200
}
197-
d := make([][]int, len(s)+1)
198-
for i := range d {
199-
d[i] = make([]int, len(t)+1)
200-
d[i][0] = i
201+
// The distance is symmetric, so iterate with the shorter string along the
202+
// row to keep the allocation as small as possible.
203+
if len(t) > len(s) {
204+
s, t = t, s
201205
}
202-
for j := range d[0] {
203-
d[0][j] = j
206+
// row[j] holds d[i][j] once column j of the current row has been written,
207+
// and d[i-1][j] until then.
208+
row := make([]int, len(t)+1)
209+
for j := range row {
210+
row[j] = j
204211
}
205-
for j := 1; j <= len(t); j++ {
206-
for i := 1; i <= len(s); i++ {
212+
var diag, above int
213+
for i := 1; i <= len(s); i++ {
214+
// diag carries d[i-1][j-1], which row[j-1] is about to overwrite.
215+
diag = row[0]
216+
row[0] = i
217+
for j := 1; j <= len(t); j++ {
218+
// d[i-1][j], the next iteration's diagonal
219+
above = row[j]
207220
if s[i-1] == t[j-1] {
208-
d[i][j] = d[i-1][j-1]
221+
row[j] = diag
209222
} else {
210-
min := d[i-1][j]
211-
if d[i][j-1] < min {
212-
min = d[i][j-1]
223+
min := above
224+
if row[j-1] < min {
225+
min = row[j-1]
213226
}
214-
if d[i-1][j-1] < min {
215-
min = d[i-1][j-1]
227+
if diag < min {
228+
min = diag
216229
}
217-
d[i][j] = min + 1
230+
row[j] = min + 1
218231
}
232+
diag = above
219233
}
220-
221234
}
222-
return d[len(s)][len(t)]
235+
return row[len(t)]
223236
}
224237

225238
func stringInSlice(a string, list []string) bool {

‎cobra_test.go‎

Lines changed: 221 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -15,6 +15,7 @@
1515
package cobra
1616

1717
import (
18+
"math/rand"
1819
"os"
1920
"os/exec"
2021
"path/filepath"
@@ -299,3 +300,223 @@ func main() {
299300
t.Error("compiled programs contains MethodByName symbol")
300301
}
301302
}
303+
304+
// ldReference is the straightforward full-matrix levenshtein implementation.
305+
// It is intentionally kept as dumb as possible so that it can be trusted as
306+
// the oracle for the single-row implementation used by ld.
307+
func ldReference(s, t string, ignoreCase bool) int {
308+
if ignoreCase {
309+
s = strings.ToLower(s)
310+
t = strings.ToLower(t)
311+
}
312+
d := make([][]int, len(s)+1)
313+
for i := range d {
314+
d[i] = make([]int, len(t)+1)
315+
d[i][0] = i
316+
}
317+
for j := range d[0] {
318+
d[0][j] = j
319+
}
320+
for j := 1; j <= len(t); j++ {
321+
for i := 1; i <= len(s); i++ {
322+
if s[i-1] == t[j-1] {
323+
d[i][j] = d[i-1][j-1]
324+
} else {
325+
min := d[i-1][j]
326+
if d[i][j-1] < min {
327+
min = d[i][j-1]
328+
}
329+
if d[i-1][j-1] < min {
330+
min = d[i-1][j-1]
331+
}
332+
d[i][j] = min + 1
333+
}
334+
}
335+
}
336+
return d[len(s)][len(t)]
337+
}
338+
339+
func TestLdBasicCases(t *testing.T) {
340+
tests := []struct {
341+
name string
342+
s string
343+
t string
344+
ignoreCase bool
345+
want int
346+
}{
347+
{
348+
name: "both empty",
349+
s: "",
350+
t: "",
351+
want: 0,
352+
},
353+
{
354+
name: "empty source",
355+
s: "",
356+
t: "abc",
357+
want: 3,
358+
},
359+
{
360+
name: "empty target",
361+
s: "abc",
362+
t: "",
363+
want: 3,
364+
},
365+
{
366+
name: "same string",
367+
s: "hello",
368+
t: "hello",
369+
want: 0,
370+
},
371+
{
372+
name: "single replace",
373+
s: "a",
374+
t: "b",
375+
want: 1,
376+
},
377+
{
378+
name: "insert",
379+
s: "abc",
380+
t: "abcd",
381+
want: 1,
382+
},
383+
{
384+
name: "delete",
385+
s: "abcd",
386+
t: "abc",
387+
want: 1,
388+
},
389+
{
390+
name: "replace",
391+
s: "kitten",
392+
t: "sitten",
393+
want: 1,
394+
},
395+
{
396+
name: "swap to shorter row",
397+
s: "a",
398+
t: "abcdefgh",
399+
want: 7,
400+
},
401+
{
402+
name: "longer source",
403+
s: "abcdefgh",
404+
t: "a",
405+
want: 7,
406+
},
407+
{
408+
name: "ignore case",
409+
s: "Hello",
410+
t: "hello",
411+
ignoreCase: true,
412+
want: 0,
413+
},
414+
{
415+
name: "multi byte string",
416+
s: "café",
417+
t: "cafe",
418+
want: 2,
419+
},
420+
}
421+
422+
for _, tt := range tests {
423+
t.Run(tt.name, func(t *testing.T) {
424+
got := ld(tt.s, tt.t, tt.ignoreCase)
425+
if got != tt.want {
426+
t.Fatalf("ld(%q, %q, %v) = %v, want %v",
427+
tt.s, tt.t, tt.ignoreCase, got, tt.want)
428+
}
429+
})
430+
}
431+
}
432+
433+
func TestLdMatchesReference(t *testing.T) {
434+
alphabet := []string{"a", "A", "b", "B", "c", "-", "é"}
435+
436+
// Fixed seed: any failure reported below is reproducible as-is.
437+
rng := rand.New(rand.NewSource(1))
438+
randString := func(n int) string {
439+
var sb strings.Builder
440+
for i := 0; i < n; i++ {
441+
sb.WriteString(alphabet[rng.Intn(len(alphabet))])
442+
}
443+
return sb.String()
444+
}
445+
446+
for i := 0; i < 20000; i++ {
447+
s := randString(rng.Intn(13))
448+
u := randString(rng.Intn(13))
449+
450+
for _, ignoreCase := range []bool{false, true} {
451+
want := ldReference(s, u, ignoreCase)
452+
got := ld(s, u, ignoreCase)
453+
if got != want {
454+
t.Fatalf("ld(%q, %q, %v) = %v, want %v", s, u, ignoreCase, got, want)
455+
}
456+
457+
// The distance is symmetric; ld swaps its arguments internally, so
458+
// assert the caller cannot observe that.
459+
if rev := ld(u, s, ignoreCase); rev != want {
460+
t.Fatalf("ld(%q, %q, %v) = %v, want %v (asymmetric)", u, s, ignoreCase, rev, want)
461+
}
462+
463+
// Distance is bounded by the length difference from below and by
464+
// the longer string from above.
465+
lo, hi := len(s)-len(u), len(s)
466+
if lo < 0 {
467+
lo = -lo
468+
}
469+
if len(u) > hi {
470+
hi = len(u)
471+
}
472+
// Bounds hold for the byte lengths only when no case folding
473+
// changed them.
474+
if !ignoreCase && (got < lo || got > hi) {
475+
t.Fatalf("ld(%q, %q, false) = %v, want within [%v, %v]", s, u, got, lo, hi)
476+
}
477+
}
478+
479+
if d := ld(s, s, false); d != 0 {
480+
t.Fatalf("ld(%q, %q, false) = %v, want 0", s, s, d)
481+
}
482+
}
483+
}
484+
485+
// ldSink keeps the compiler from optimizing the benchmarked calls away.
486+
var ldSink int
487+
488+
// BenchmarkLd compares the single-row implementation against the full-matrix
489+
// one over the input sizes cobra actually sees: SuggestionsFor measures a
490+
// mistyped command name against every sibling command name, so both operands
491+
// are command names rather than arbitrary user input. The "long" case is well
492+
// past anything realistic and is only there to show how the two implementations
493+
// diverge as the inputs grow.
494+
func BenchmarkLd(b *testing.B) {
495+
benchmarks := []struct {
496+
name string
497+
s, t string
498+
}{
499+
{"short", "get", "set"},
500+
{"typical", "kubectl-config", "kubectl-configs"},
501+
{"long", strings.Repeat("abcde-", 8), strings.Repeat("abdce-", 8)},
502+
}
503+
504+
impls := []struct {
505+
name string
506+
fn func(s, t string, ignoreCase bool) int
507+
}{
508+
{"optimized", ld},
509+
{"reference", ldReference},
510+
}
511+
512+
for _, bm := range benchmarks {
513+
for _, impl := range impls {
514+
b.Run(bm.name+"/"+impl.name, func(b *testing.B) {
515+
b.ReportAllocs()
516+
for i := 0; i < b.N; i++ {
517+
ldSink = impl.fn(bm.s, bm.t, true)
518+
}
519+
})
520+
}
521+
}
522+
}

0 commit comments

Comments
 (0)