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:
@@ -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"))
|
||||
|
||||
@@ -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},
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
@@ -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,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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]
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user