feat(scoring): transition workflow learning and fix history UX (#39)

- add transition scoring engine to learn sequential workflows per
directory using command skeletons
- prioritize current active git branch over older branches in
suggestions
- restore strict chronological order for history navigation and bypass
AI re-ranking
- add in-memory session history for instant access to just-run commands
- fix PTY prompt not syncing when using arrow keys in the history menu
- fix history tie-breaker to properly prioritize newer commands by
assigning larger IDs
- fix potential mutex deadlock between bufferMu and pty write during
history navigation
- fix workspace git detection to remove depth limits and correctly
identify repositories with detached HEADs
- fix SQLite connection leaks in frecency queries
- fix global state leaks in history tests by snapshotting and restoring
registry states
- modernize string prefix checks using strings.CutPrefix
This commit is contained in:
VERSE
2026-07-26 19:10:56 +07:00
committed by GitHub
parent 85ac76e85a
commit 6ffb9f83f2
20 changed files with 1004 additions and 72 deletions
+16
View File
@@ -36,6 +36,22 @@ var DefaultContextRules = []ContextRule{
},
bonus: 40,
},
&SimpleContextRule{
check: func(ws workspace.WorkspaceInfo, cmd string) bool {
if !ws.HasGit || ws.GitBranch == "" {
return false
}
isRelevantCmd := strings.HasPrefix(cmd, "git push") || strings.HasPrefix(cmd, "git pull") ||
strings.HasPrefix(cmd, "git checkout") || strings.HasPrefix(cmd, "git switch") ||
strings.HasPrefix(cmd, "git branch") || strings.HasPrefix(cmd, "git merge") ||
strings.HasPrefix(cmd, "git rebase")
if !isRelevantCmd {
return false
}
return strings.HasSuffix(cmd, " "+ws.GitBranch) || strings.Contains(cmd, " "+ws.GitBranch+" ")
},
bonus: 60,
},
&SimpleContextRule{
check: func(ws workspace.WorkspaceInfo, cmd string) bool {
return ws.HasGit && (strings.HasPrefix(cmd, "git init") || strings.HasPrefix(cmd, "git clone"))
+12
View File
@@ -61,6 +61,18 @@ func TestApplyContextRules_MultiEcosystem(t *testing.T) {
cmd: "kubectl get pods",
expected: 40,
},
{
name: "git push active branch gets double bonus clamped to 100",
ws: workspace.WorkspaceInfo{HasGit: true, GitBranch: "fix/scoring"},
cmd: "git push origin fix/scoring",
expected: 100, // 40 (git push) + 60 (active branch) = 100
},
{
name: "git push other branch gets normal git push bonus only",
ws: workspace.WorkspaceInfo{HasGit: true, GitBranch: "fix/scoring"},
cmd: "git push origin fix/alias",
expected: 40,
},
{
name: "unrelated command gets no bonus",
ws: workspace.WorkspaceInfo{HasGit: true},
+172 -2
View File
@@ -22,6 +22,14 @@ type FrecencyEntry struct {
RawScore float64
}
type TransitionEntry struct {
PrevSkeleton string
NextSkeleton string
Cwd string
Count int
LastUsed time.Time
}
type FrecencyStore struct {
db *sql.DB
mu sync.Mutex
@@ -51,6 +59,7 @@ func NewFrecencyStore(dbPath string) (*FrecencyStore, error) {
if err != nil {
return nil, fmt.Errorf("failed to open sqlite database: %w", err)
}
db.SetMaxOpenConns(1)
store := &FrecencyStore{db: db}
if err := store.initSchema(context.Background()); err != nil {
@@ -89,16 +98,29 @@ CREATE TABLE IF NOT EXISTS history_entries (
);
CREATE INDEX IF NOT EXISTS idx_history_cwd_cmd ON history_entries(cwd, cmd);
CREATE TABLE IF NOT EXISTS command_transitions (
id INTEGER PRIMARY KEY AUTOINCREMENT,
prev_skeleton TEXT NOT NULL,
next_skeleton TEXT NOT NULL,
cwd TEXT NOT NULL,
count INTEGER DEFAULT 1,
last_used TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
UNIQUE(prev_skeleton, next_skeleton, cwd)
);
CREATE INDEX IF NOT EXISTS idx_transitions_prev_cwd ON command_transitions(prev_skeleton, cwd);
`
_, err := f.db.ExecContext(ctxTimeout, schema)
return err
}
func (f *FrecencyStore) Record(ctx context.Context, cmd, cwd string) error {
func (f *FrecencyStore) Record(ctx context.Context, cmd, cwd string, exitCode int) error {
if f == nil {
return nil
}
cmd = strings.TrimSpace(cmd)
cwd = strings.TrimSpace(cwd)
if cmd == "" || cwd == "" {
return nil
}
@@ -112,17 +134,165 @@ func (f *FrecencyStore) Record(ctx context.Context, cmd, cwd string) error {
ctxTimeout, cancel := context.WithTimeout(ctx, 1000*time.Millisecond)
defer cancel()
query := `
var query string
if exitCode == 0 {
query = `
INSERT INTO history_entries (cmd, cwd, count, last_used)
VALUES (?, ?, 1, CURRENT_TIMESTAMP)
ON CONFLICT(cmd, cwd) DO UPDATE SET
count = count + 1,
last_used = CURRENT_TIMESTAMP;
`
} else {
query = `
INSERT INTO history_entries (cmd, cwd, count, last_used)
VALUES (?, ?, 0, CURRENT_TIMESTAMP)
ON CONFLICT(cmd, cwd) DO UPDATE SET
last_used = CURRENT_TIMESTAMP;
`
}
_, err := f.db.ExecContext(ctxTimeout, query, cmd, cwd)
return err
}
func (f *FrecencyStore) RecordTransition(ctx context.Context, prevSkeleton, nextSkeleton, cwd string, nextExitCode int) error {
if f == nil {
return nil
}
prevSkeleton = strings.TrimSpace(prevSkeleton)
nextSkeleton = strings.TrimSpace(nextSkeleton)
cwd = strings.TrimSpace(cwd)
if prevSkeleton == "" || nextSkeleton == "" || cwd == "" {
return nil
}
f.mu.Lock()
defer f.mu.Unlock()
if ctx == nil {
ctx = context.Background()
}
ctxTimeout, cancel := context.WithTimeout(ctx, 1000*time.Millisecond)
defer cancel()
var query string
if nextExitCode == 0 {
query = `
INSERT INTO command_transitions (prev_skeleton, next_skeleton, cwd, count, last_used)
VALUES (?, ?, ?, 1, CURRENT_TIMESTAMP)
ON CONFLICT(prev_skeleton, next_skeleton, cwd) DO UPDATE SET
count = count + 1,
last_used = CURRENT_TIMESTAMP;
`
} else {
query = `
INSERT INTO command_transitions (prev_skeleton, next_skeleton, cwd, count, last_used)
VALUES (?, ?, ?, 0, CURRENT_TIMESTAMP)
ON CONFLICT(prev_skeleton, next_skeleton, cwd) DO UPDATE SET
last_used = CURRENT_TIMESTAMP;
`
}
_, err := f.db.ExecContext(ctxTimeout, query, prevSkeleton, nextSkeleton, cwd)
return err
}
func (f *FrecencyStore) QueryTransitionsWithFallback(ctx context.Context, prevSkeleton, cwd string) ([]TransitionEntry, bool) {
if f == nil {
return nil, false
}
prevSkeleton = strings.TrimSpace(prevSkeleton)
cwd = strings.TrimSpace(cwd)
if prevSkeleton == "" {
return nil, false
}
f.mu.Lock()
defer f.mu.Unlock()
if ctx == nil {
ctx = context.Background()
}
ctxTimeout, cancel := context.WithTimeout(ctx, 1000*time.Millisecond)
defer cancel()
// Phase 1: Local query with depth fallback
parts := strings.Fields(prevSkeleton)
for len(parts) > 0 {
key := strings.Join(parts, " ")
var loopEntries []TransitionEntry
func() {
rows, err := f.db.QueryContext(ctxTimeout, `
SELECT prev_skeleton, next_skeleton, cwd, count, last_used
FROM command_transitions
WHERE prev_skeleton = ? AND cwd = ? AND count > 0
ORDER BY count DESC
`, key, cwd)
if err == nil {
defer rows.Close()
for rows.Next() {
var prev, next, rCwd string
var count int
var lastUsedRaw string
if err := rows.Scan(&prev, &next, &rCwd, &count, &lastUsedRaw); err == nil {
t, _ := parseTimestamp(lastUsedRaw)
loopEntries = append(loopEntries, TransitionEntry{
PrevSkeleton: prev,
NextSkeleton: next,
Cwd: rCwd,
Count: count,
LastUsed: t,
})
}
}
}
}()
if len(loopEntries) > 0 {
return loopEntries, true
}
parts = parts[:len(parts)-1]
}
// Phase 2: Global query with depth fallback
parts = strings.Fields(prevSkeleton)
for len(parts) > 0 {
key := strings.Join(parts, " ")
var loopEntries []TransitionEntry
func() {
rows, err := f.db.QueryContext(ctxTimeout, `
SELECT prev_skeleton, next_skeleton, SUM(count) as total_count, MAX(last_used) as max_last_used
FROM command_transitions
WHERE prev_skeleton = ? AND count > 0
GROUP BY next_skeleton
ORDER BY total_count DESC
`, key)
if err == nil {
defer rows.Close()
for rows.Next() {
var prev, next string
var count int
var lastUsedRaw string
if err := rows.Scan(&prev, &next, &count, &lastUsedRaw); err == nil {
t, _ := parseTimestamp(lastUsedRaw)
loopEntries = append(loopEntries, TransitionEntry{
PrevSkeleton: prev,
NextSkeleton: next,
Cwd: "",
Count: count,
LastUsed: t,
})
}
}
}
}()
if len(loopEntries) > 0 {
return loopEntries, false
}
parts = parts[:len(parts)-1]
}
return nil, false
}
func (f *FrecencyStore) RawScore(count int, lastUsed time.Time) float64 {
if count <= 0 {
return 0
+76 -9
View File
@@ -19,10 +19,10 @@ func TestFrecencyStore_RecordAndQueryLocal(t *testing.T) {
defer store.Close()
cwd := "/home/user/project"
_ = store.Record(context.Background(), "git status", cwd)
_ = store.Record(context.Background(), "git status", cwd)
_ = store.Record(context.Background(), "git status", cwd)
_ = store.Record(context.Background(), "git commit -m 'test'", cwd)
_ = store.Record(context.Background(), "git status", cwd, 0)
_ = store.Record(context.Background(), "git status", cwd, 0)
_ = store.Record(context.Background(), "git status", cwd, 0)
_ = store.Record(context.Background(), "git commit -m 'test'", cwd, 0)
entries, err := store.QueryLocal(context.Background(), cwd, "git", 10)
if err != nil {
@@ -60,9 +60,9 @@ func TestFrecencyStore_QueryGlobalDedupe(t *testing.T) {
}
defer store.Close()
_ = store.Record(context.Background(), "make build", "/repo/a")
_ = store.Record(context.Background(), "make build", "/repo/a")
_ = store.Record(context.Background(), "make build", "/repo/b")
_ = store.Record(context.Background(), "make build", "/repo/a", 0)
_ = store.Record(context.Background(), "make build", "/repo/a", 0)
_ = store.Record(context.Background(), "make build", "/repo/b", 0)
entries, err := store.QueryGlobal(context.Background(), "make", 10)
if err != nil {
@@ -139,7 +139,7 @@ func TestFrecencyStore_SQLiteConfigurationAndContext(t *testing.T) {
ctxCanceled, cancel := context.WithCancel(context.Background())
cancel()
err = store.Record(ctxCanceled, "git status", tmpDir)
err = store.Record(ctxCanceled, "git status", tmpDir, 0)
if !errors.Is(err, context.Canceled) {
t.Errorf("expected context.Canceled from Record with canceled context, got %v", err)
}
@@ -147,7 +147,7 @@ func TestFrecencyStore_SQLiteConfigurationAndContext(t *testing.T) {
func TestFrecencyStore_NilReceiver(t *testing.T) {
var nilStore *FrecencyStore
if err := nilStore.Record(context.Background(), "cmd", "cwd"); err != nil {
if err := nilStore.Record(context.Background(), "cmd", "cwd", 0); err != nil {
t.Errorf("expected nil error on nil store Record, got %v", err)
}
if entries, err := nilStore.QueryLocal(context.Background(), "cwd", "", 10); err != nil || entries != nil {
@@ -160,3 +160,70 @@ func TestFrecencyStore_NilReceiver(t *testing.T) {
t.Errorf("expected nil error on nil store Close, got %v", err)
}
}
func TestFrecencyStore_ExitCodeBehavior(t *testing.T) {
tmpDir := t.TempDir()
dbPath := filepath.Join(tmpDir, "history.db")
store, err := NewFrecencyStore(dbPath)
if err != nil {
t.Fatalf("NewFrecencyStore failed: %v", err)
}
defer store.Close()
cwd := "/home/user/test"
_ = store.Record(context.Background(), "grep foo", cwd, 0) // count=1
_ = store.Record(context.Background(), "grep foo", cwd, 1) // count unchanged (1)
entries, _ := store.QueryLocal(context.Background(), cwd, "grep", 10)
if len(entries) != 1 || entries[0].Count != 1 {
t.Errorf("expected grep count to be 1 after non-zero exit code, got %v", entries)
}
_ = store.RecordTransition(context.Background(), "git checkout", "git status", cwd, 0)
_ = store.RecordTransition(context.Background(), "git checkout", "git status", cwd, 1)
transitions, isLocal := store.QueryTransitionsWithFallback(context.Background(), "git checkout", cwd)
if !isLocal || len(transitions) != 1 || transitions[0].Count != 1 {
t.Errorf("expected transition count 1 after non-zero exit code, got %v, isLocal=%v", transitions, isLocal)
}
}
func TestFrecencyStore_TransitionCwdIsolationAndDepthFallback(t *testing.T) {
tmpDir := t.TempDir()
dbPath := filepath.Join(tmpDir, "history.db")
store, err := NewFrecencyStore(dbPath)
if err != nil {
t.Fatalf("NewFrecencyStore failed: %v", err)
}
defer store.Close()
projectA := "/repo/a"
projectB := "/repo/b"
_ = store.RecordTransition(context.Background(), "git checkout", "npm run dev", projectA, 0)
_ = store.RecordTransition(context.Background(), "git checkout", "go test", projectB, 0)
_ = store.RecordTransition(context.Background(), "git checkout", "go test", projectB, 0)
// query in project B should return go test (Local) and not npm run dev
transB, isLocalB := store.QueryTransitionsWithFallback(context.Background(), "git checkout", projectB)
if !isLocalB || len(transB) != 1 || transB[0].NextSkeleton != "go test" {
t.Errorf("expected local transition 'go test' for project B, got %v (isLocal=%v)", transB, isLocalB)
}
// query in project C (no local data) should fallback to Global (returning both aggregated)
projectC := "/repo/c"
transC, isLocalC := store.QueryTransitionsWithFallback(context.Background(), "git checkout", projectC)
if isLocalC || len(transC) != 2 {
t.Errorf("expected global transitions for project C, got %v (isLocal=%v)", transC, isLocalC)
}
if transC[0].NextSkeleton != "go test" {
t.Errorf("expected global top transition to be 'go test' (count 2), got %s", transC[0].NextSkeleton)
}
// depth fallback test: query deep skeleton with no exact match should fallback to shallower prefix
_ = store.RecordTransition(context.Background(), "git remote", "git fetch", projectA, 0)
transDeep, isLocalDeep := store.QueryTransitionsWithFallback(context.Background(), "git remote add", projectA)
if !isLocalDeep || len(transDeep) != 1 || transDeep[0].NextSkeleton != "git fetch" {
t.Errorf("expected depth fallback to 'git fetch' from 'git remote', got %v", transDeep)
}
}
+30 -1
View File
@@ -12,6 +12,7 @@ type ScoreBreakdown struct {
BasePriority int
ContextBonus int
Frecency int
Transition int
MatchQuality int
}
@@ -25,13 +26,15 @@ type ScoreConfig struct {
WeightBasePriority float64
WeightContextBonus float64
WeightFrecency float64
WeightTransition float64
WeightMatchQuality float64
}
var DefaultScoreConfig = ScoreConfig{
WeightBasePriority: 0.30,
WeightContextBonus: 0.25,
WeightFrecency: 0.25,
WeightFrecency: 0.15,
WeightTransition: 0.10,
WeightMatchQuality: 0.20,
}
@@ -71,11 +74,13 @@ func ScoreWithConfig(suggestions []spec.Suggestion, signals SignalSet, config Sc
bp := basePriorityFor(s)
cb := ApplyContextRules(signals.Workspace, s.Cmd)
frec := normFrec[i]
trans := transitionScoreFor(ExtractSkeleton(s.Cmd), signals.TransitionEntries, signals.TransitionIsLocal)
mq := matchQualityScore(s.Cmd, signals.Query)
total := config.WeightBasePriority*float64(bp) +
config.WeightContextBonus*float64(cb) +
config.WeightFrecency*float64(frec) +
config.WeightTransition*float64(trans) +
config.WeightMatchQuality*float64(mq)
scored[i] = ScoredSuggestion{
@@ -85,6 +90,7 @@ func ScoreWithConfig(suggestions []spec.Suggestion, signals SignalSet, config Sc
BasePriority: bp,
ContextBonus: cb,
Frecency: frec,
Transition: trans,
MatchQuality: mq,
},
}
@@ -94,6 +100,9 @@ func ScoreWithConfig(suggestions []spec.Suggestion, signals SignalSet, config Sc
if scored[i].Score != scored[j].Score {
return scored[i].Score > scored[j].Score
}
if scored[i].Breakdown.Transition != scored[j].Breakdown.Transition {
return scored[i].Breakdown.Transition > scored[j].Breakdown.Transition
}
if scored[i].Breakdown.Frecency != scored[j].Breakdown.Frecency {
return scored[i].Breakdown.Frecency > scored[j].Breakdown.Frecency
}
@@ -106,6 +115,26 @@ func ScoreWithConfig(suggestions []spec.Suggestion, signals SignalSet, config Sc
return scored
}
func transitionScoreFor(cmdSkeleton string, entries []TransitionEntry, isLocal bool) int {
if len(entries) == 0 {
return 0 // cold-start: no data, contributes 0 (must check before accessing entries[0])
}
maxCount := entries[0].Count
if maxCount <= 0 {
return 0
}
for _, e := range entries {
if e.NextSkeleton == cmdSkeleton {
score := (float64(e.Count) / float64(maxCount)) * 100.0
if !isLocal {
score *= 0.7
}
return int(math.Round(score))
}
}
return 0
}
func basePriorityFor(s spec.Suggestion) int {
if s.Priority > 0 {
if s.Priority > 100 {
+80
View File
@@ -156,3 +156,83 @@ func TestIsSubsequence_UTF8(t *testing.T) {
t.Error("expected multi-byte rune 'αγ' to be subsequence of 'αβγδε'")
}
}
func TestScore_TransitionAndColdStartGuard(t *testing.T) {
spec.ResetRegistry()
spec.Register(&spec.Spec{
Name: "git",
Subcommands: []spec.Subcommand{
{Name: "pull"},
{Name: "status"},
},
})
suggestions := []spec.Suggestion{
{Cmd: "git pull --rebase origin main", Source: "spec"},
{Cmd: "git status", Source: "spec"},
}
// Cold-start test: nil/empty TransitionEntries should not panic and contribute 0
signalsCold := SignalSet{
Query: "git",
TransitionEntries: nil,
}
scoredCold := Score(suggestions, signalsCold)
if len(scoredCold) != 2 || scoredCold[0].Breakdown.Transition != 0 {
t.Errorf("expected 0 transition score on cold-start without panic, got %v", scoredCold)
}
// Transition Local test
signalsLocal := SignalSet{
Query: "git",
TransitionEntries: []TransitionEntry{{NextSkeleton: "git pull", Count: 10}},
TransitionIsLocal: true,
}
scoredLocal := Score(suggestions, signalsLocal)
if scoredLocal[0].Breakdown.Transition != 100 {
t.Errorf("expected transition score 100 for git pull when local, got %d", scoredLocal[0].Breakdown.Transition)
}
// Transition Global test (damping 70%)
signalsGlobal := SignalSet{
Query: "git",
TransitionEntries: []TransitionEntry{{NextSkeleton: "git pull", Count: 10}},
TransitionIsLocal: false,
}
scoredGlobal := Score(suggestions, signalsGlobal)
if scoredGlobal[0].Breakdown.Transition != 70 {
t.Errorf("expected transition score 70 for git pull when global (70%% damping), got %d", scoredGlobal[0].Breakdown.Transition)
}
}
func TestScore_TieBreakingOrder(t *testing.T) {
// All three suggestions will be constructed to have identical total scores (0),
// but different breakdown scores to verify tie-break priority: Transition > Frecency > ContextBonus > Alphabetical.
suggestions := []spec.Suggestion{
{Cmd: "cmdA", Source: "test", Priority: 0},
{Cmd: "cmdB", Source: "test", Priority: 0},
}
// We set weight configuration with 0 weights so total scores are identical (0), triggering tie-break logic
zeroConfig := ScoreConfig{
WeightBasePriority: 0,
WeightContextBonus: 0,
WeightFrecency: 0,
WeightTransition: 0,
WeightMatchQuality: 0,
}
signals := SignalSet{
Query: "",
TransitionEntries: []TransitionEntry{
{NextSkeleton: "cmdB", Count: 10},
{NextSkeleton: "cmdA", Count: 5},
},
TransitionIsLocal: true,
}
scored := ScoreWithConfig(suggestions, signals, zeroConfig)
if len(scored) != 2 || scored[0].Cmd != "cmdB" {
t.Errorf("expected cmdB to win tie-break due to higher Transition score, got %v", scored)
}
}
+24 -14
View File
@@ -8,16 +8,18 @@ import (
)
type SignalSet struct {
Workspace workspace.WorkspaceInfo
LocalFrecency []FrecencyEntry
GlobalFrecency []FrecencyEntry
Query string
RootCommand string
Cwd string
Workspace workspace.WorkspaceInfo
LocalFrecency []FrecencyEntry
GlobalFrecency []FrecencyEntry
TransitionEntries []TransitionEntry
TransitionIsLocal bool
Query string
RootCommand string
Cwd string
}
// CollectSignals gathers environment, workspace, and historical frecency signals for the given query and directory
func CollectSignals(ctx context.Context, cwd, query, rootCmd string, frecency *FrecencyStore) SignalSet {
// CollectSignals gathers environment, workspace, and historical frecency/transition signals for the given query and directory
func CollectSignals(ctx context.Context, cwd, query, rootCmd string, frecency *FrecencyStore, prevCmdSkeleton string) SignalSet {
ws := workspace.DetectCached(cwd)
if ctx == nil {
@@ -25,17 +27,25 @@ func CollectSignals(ctx context.Context, cwd, query, rootCmd string, frecency *F
}
var local, global []FrecencyEntry
var trans []TransitionEntry
var transIsLocal bool
if frecency != nil {
local, _ = frecency.QueryLocal(ctx, cwd, query, 50)
global, _ = frecency.QueryGlobal(ctx, query, 50)
if prevCmdSkeleton != "" {
trans, transIsLocal = frecency.QueryTransitionsWithFallback(ctx, prevCmdSkeleton, cwd)
}
}
return SignalSet{
Workspace: ws,
LocalFrecency: local,
GlobalFrecency: global,
Query: strings.TrimSpace(query),
RootCommand: strings.TrimSpace(rootCmd),
Cwd: cwd,
Workspace: ws,
LocalFrecency: local,
GlobalFrecency: global,
TransitionEntries: trans,
TransitionIsLocal: transIsLocal,
Query: strings.TrimSpace(query),
RootCommand: strings.TrimSpace(rootCmd),
Cwd: cwd,
}
}
+7 -3
View File
@@ -18,10 +18,11 @@ func TestCollectSignals(t *testing.T) {
}
defer store.Close()
_ = store.Record(context.Background(), "npm run dev", tmpDir)
_ = store.Record(context.Background(), "npm test", "/other/dir")
_ = store.Record(context.Background(), "npm run dev", tmpDir, 0)
_ = store.Record(context.Background(), "npm test", "/other/dir", 0)
_ = store.RecordTransition(context.Background(), "git checkout", "npm run dev", tmpDir, 0)
signals := CollectSignals(context.Background(), tmpDir, "npm", "npm", store)
signals := CollectSignals(context.Background(), tmpDir, "npm", "npm", store, "git checkout")
if !signals.Workspace.HasNodeProject {
t.Error("expected HasNodeProject to be true in collected signals")
@@ -32,4 +33,7 @@ func TestCollectSignals(t *testing.T) {
if len(signals.GlobalFrecency) != 2 {
t.Errorf("expected global frecency to contain 2 entries, got %d", len(signals.GlobalFrecency))
}
if !signals.TransitionIsLocal || len(signals.TransitionEntries) != 1 || signals.TransitionEntries[0].NextSkeleton != "npm run dev" {
t.Errorf("expected transition entry 'npm run dev', got %v (isLocal=%v)", signals.TransitionEntries, signals.TransitionIsLocal)
}
}
+23
View File
@@ -0,0 +1,23 @@
package scoring
import (
"strings"
"github.com/versenilvis/iris/spec"
)
// ExtractSkeleton extracts the subcommand skeleton of a command string for transition tracking.
// If a spec is registered, it walks the subcommand tree.
// If no spec is registered (or fallback occurs), it returns the first token (binary name).
// Never returns empty string unless input has no tokens.
func ExtractSkeleton(buf string) string {
if skeleton, ok := spec.TryExtractSkeleton(buf); ok && skeleton != "" {
return skeleton
}
fields := strings.Fields(buf)
if len(fields) == 0 {
return ""
}
return fields[0]
}
+54
View File
@@ -0,0 +1,54 @@
package scoring
import (
"testing"
"github.com/versenilvis/iris/spec"
)
func TestExtractSkeleton(t *testing.T) {
spec.ResetRegistry()
spec.Register(&spec.Spec{
Name: "git",
Subcommands: []spec.Subcommand{
{Name: "checkout"},
{Name: "push"},
},
})
tests := []struct {
name string
buf string
want string
}{
{
name: "spec command with args",
buf: "git checkout feature-x",
want: "git checkout",
},
{
name: "fallback no spec binary only",
buf: "cargo build --release",
want: "cargo",
},
{
name: "fallback single command",
buf: "ls -la",
want: "ls",
},
{
name: "empty input",
buf: " ",
want: "",
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
got := ExtractSkeleton(tc.buf)
if got != tc.want {
t.Errorf("ExtractSkeleton(%q) = %q, want %q", tc.buf, got, tc.want)
}
})
}
}