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
+7 -2
View File
@@ -6,6 +6,7 @@ import (
"strings" "strings"
"time" "time"
"github.com/versenilvis/iris/internal/workspace"
"github.com/versenilvis/iris/spec" "github.com/versenilvis/iris/spec"
) )
@@ -67,8 +68,8 @@ func getGitResultsFiltered(tokens []string, localOnly bool, args ...string) []sp
return nil return nil
} }
activeBranch := "" activeBranch := workspace.DetectCached(cwd).GitBranch
if args[0] == "branch" { if activeBranch == "" && args[0] == "branch" {
activeCmd := exec.CommandContext(ctx, "git", "rev-parse", "--abbrev-ref", "HEAD") activeCmd := exec.CommandContext(ctx, "git", "rev-parse", "--abbrev-ref", "HEAD")
activeCmd.Dir = cwd activeCmd.Dir = cwd
if activeOut, err := activeCmd.Output(); err == nil { if activeOut, err := activeCmd.Output(); err == nil {
@@ -152,9 +153,13 @@ func getGitResultsFiltered(tokens []string, localOnly bool, args ...string) []sp
}) })
} }
if activeBranch == "" {
activeBranch = workspace.DetectCached(cwd).GitBranch
}
if activeBranch != "" { if activeBranch != "" {
for i, r := range results { for i, r := range results {
if r.Cmd == activeBranch { if r.Cmd == activeBranch {
r.Priority = 100
copy(results[1:i+1], results[0:i]) copy(results[1:i+1], results[0:i])
results[0] = r results[0] = r
break break
+75 -33
View File
@@ -13,6 +13,9 @@ import (
) )
var ( var (
sessionHistory []string
sessionHistoryMu sync.Mutex
historyCache []string historyCache []string
idMapCache map[string]int idMapCache map[string]int
searcherCache *fuzzy.Searcher searcherCache *fuzzy.Searcher
@@ -20,6 +23,24 @@ var (
lastModTime int64 lastModTime int64
) )
func RecordSessionCommand(cmd string) {
cmd = strings.TrimSpace(cmd)
if cmd == "" {
return
}
mu.Lock()
defer mu.Unlock()
sessionHistoryMu.Lock()
defer sessionHistoryMu.Unlock()
if len(sessionHistory) > 0 && sessionHistory[len(sessionHistory)-1] == cmd {
return
}
sessionHistory = append(sessionHistory, cmd)
historyCache = nil // invalidate to merge session history on next search
}
type HistResult struct { type HistResult struct {
ID int ID int
Cmd string Cmd string
@@ -65,61 +86,82 @@ func SearchHistory(query string, aliases map[string]string) ([]HistResult, error
// lazy load history if cache is empty // lazy load history if cache is empty
if len(historyCache) == 0 { if len(historyCache) == 0 {
file, err := os.Open(histFile) file, err := os.Open(histFile)
if err != nil { if err != nil && !os.IsNotExist(err) {
return nil, err return nil, err
} }
defer func() { _ = file.Close() }() if file != nil {
defer func() { _ = file.Close() }()
}
var allCmds []string var allCmds []string
scanner := bufio.NewScanner(file) if file != nil {
for scanner.Scan() { scanner := bufio.NewScanner(file)
line := scanner.Text() for scanner.Scan() {
cmd := line line := scanner.Text()
cmd := line
if shellName == "zsh" { if shellName == "zsh" {
parts := strings.SplitN(line, ";", 2) parts := strings.SplitN(line, ";", 2)
if len(parts) == 2 { if len(parts) == 2 {
cmd = parts[1] cmd = parts[1]
} }
} else if shellName == "bash" { } else if shellName == "bash" {
if strings.HasPrefix(line, "#") && len(line) > 1 { if strings.HasPrefix(line, "#") && len(line) > 1 {
isTimestamp := true isTimestamp := true
for _, c := range line[1:] { for _, c := range line[1:] {
if c < '0' || c > '9' { if c < '0' || c > '9' {
isTimestamp = false isTimestamp = false
break break
}
}
if isTimestamp {
continue
} }
} }
if isTimestamp { } else if shellName == "fish" {
if after, ok := strings.CutPrefix(line, "- cmd: "); ok {
cmd = after
} else {
continue continue
} }
} }
} else if shellName == "fish" {
if after, ok := strings.CutPrefix(line, "- cmd: "); ok { cmd = strings.TrimSpace(cmd)
cmd = after if cmd != "" {
} else { allCmds = append(allCmds, cmd)
continue
} }
} }
if err := scanner.Err(); err != nil {
cmd = strings.TrimSpace(cmd) return nil, err
if cmd != "" {
allCmds = append(allCmds, cmd)
} }
} }
if err := scanner.Err(); err != nil {
return nil, err
}
// build historyCache backwards so newest commands come first // build historyCache backwards so newest commands come first
seen := make(map[string]bool) seen := make(map[string]bool)
historyCache = nil
idMapCache = make(map[string]int)
currentID := len(sessionHistory) + len(allCmds)
sessionHistoryMu.Lock()
for i := len(sessionHistory) - 1; i >= 0; i-- {
cmd := sessionHistory[i]
if !seen[cmd] {
historyCache = append(historyCache, cmd)
seen[cmd] = true
idMapCache[cmd] = currentID
currentID--
}
}
sessionHistoryMu.Unlock()
for i := len(allCmds) - 1; i >= 0; i-- { for i := len(allCmds) - 1; i >= 0; i-- {
cmd := allCmds[i] cmd := allCmds[i]
if !seen[cmd] { if !seen[cmd] {
historyCache = append(historyCache, cmd) historyCache = append(historyCache, cmd)
seen[cmd] = true seen[cmd] = true
// we assign the ID as the original line number (1-indexed based on allCmds length) idMapCache[cmd] = currentID
idMapCache[cmd] = i + 1 currentID--
} }
} }
+52
View File
@@ -0,0 +1,52 @@
package integration
import (
"testing"
)
func TestRecordSessionCommand_MergeAndDeduplicate(t *testing.T) {
sessionHistoryMu.Lock()
origSessionHistory := sessionHistory
sessionHistory = nil
sessionHistoryMu.Unlock()
mu.Lock()
origHistoryCache := historyCache
historyCache = nil
mu.Unlock()
t.Cleanup(func() {
sessionHistoryMu.Lock()
sessionHistory = origSessionHistory
sessionHistoryMu.Unlock()
mu.Lock()
historyCache = origHistoryCache
mu.Unlock()
})
RecordSessionCommand("git status")
RecordSessionCommand("npm run dev")
RecordSessionCommand("npm run dev") // duplicate subsequent command should be ignored
RecordSessionCommand("git push origin fix/scoring")
results, err := SearchHistory("", nil)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if len(results) < 3 {
t.Fatalf("expected at least 3 session commands in search results, got %d", len(results))
}
// newest session command must be results[0]
if results[0].Cmd != "git push origin fix/scoring" {
t.Errorf("expected results[0] to be 'git push origin fix/scoring', got %q", results[0].Cmd)
}
if results[1].Cmd != "npm run dev" {
t.Errorf("expected results[1] to be 'npm run dev', got %q", results[1].Cmd)
}
if results[2].Cmd != "git status" {
t.Errorf("expected results[2] to be 'git status', got %q", results[2].Cmd)
}
}
+16
View File
@@ -36,6 +36,22 @@ var DefaultContextRules = []ContextRule{
}, },
bonus: 40, 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{ &SimpleContextRule{
check: func(ws workspace.WorkspaceInfo, cmd string) bool { check: func(ws workspace.WorkspaceInfo, cmd string) bool {
return ws.HasGit && (strings.HasPrefix(cmd, "git init") || strings.HasPrefix(cmd, "git clone")) 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", cmd: "kubectl get pods",
expected: 40, 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", name: "unrelated command gets no bonus",
ws: workspace.WorkspaceInfo{HasGit: true}, ws: workspace.WorkspaceInfo{HasGit: true},
+172 -2
View File
@@ -22,6 +22,14 @@ type FrecencyEntry struct {
RawScore float64 RawScore float64
} }
type TransitionEntry struct {
PrevSkeleton string
NextSkeleton string
Cwd string
Count int
LastUsed time.Time
}
type FrecencyStore struct { type FrecencyStore struct {
db *sql.DB db *sql.DB
mu sync.Mutex mu sync.Mutex
@@ -51,6 +59,7 @@ func NewFrecencyStore(dbPath string) (*FrecencyStore, error) {
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to open sqlite database: %w", err) return nil, fmt.Errorf("failed to open sqlite database: %w", err)
} }
db.SetMaxOpenConns(1)
store := &FrecencyStore{db: db} store := &FrecencyStore{db: db}
if err := store.initSchema(context.Background()); err != nil { 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 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) _, err := f.db.ExecContext(ctxTimeout, schema)
return err 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 { if f == nil {
return nil return nil
} }
cmd = strings.TrimSpace(cmd) cmd = strings.TrimSpace(cmd)
cwd = strings.TrimSpace(cwd)
if cmd == "" || cwd == "" { if cmd == "" || cwd == "" {
return nil 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) ctxTimeout, cancel := context.WithTimeout(ctx, 1000*time.Millisecond)
defer cancel() defer cancel()
query := ` var query string
if exitCode == 0 {
query = `
INSERT INTO history_entries (cmd, cwd, count, last_used) INSERT INTO history_entries (cmd, cwd, count, last_used)
VALUES (?, ?, 1, CURRENT_TIMESTAMP) VALUES (?, ?, 1, CURRENT_TIMESTAMP)
ON CONFLICT(cmd, cwd) DO UPDATE SET ON CONFLICT(cmd, cwd) DO UPDATE SET
count = count + 1, count = count + 1,
last_used = CURRENT_TIMESTAMP; 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) _, err := f.db.ExecContext(ctxTimeout, query, cmd, cwd)
return err 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 { func (f *FrecencyStore) RawScore(count int, lastUsed time.Time) float64 {
if count <= 0 { if count <= 0 {
return 0 return 0
+76 -9
View File
@@ -19,10 +19,10 @@ func TestFrecencyStore_RecordAndQueryLocal(t *testing.T) {
defer store.Close() defer store.Close()
cwd := "/home/user/project" cwd := "/home/user/project"
_ = store.Record(context.Background(), "git status", cwd) _ = store.Record(context.Background(), "git status", cwd, 0)
_ = store.Record(context.Background(), "git status", cwd) _ = store.Record(context.Background(), "git status", cwd, 0)
_ = store.Record(context.Background(), "git status", cwd) _ = store.Record(context.Background(), "git status", cwd, 0)
_ = store.Record(context.Background(), "git commit -m 'test'", cwd) _ = store.Record(context.Background(), "git commit -m 'test'", cwd, 0)
entries, err := store.QueryLocal(context.Background(), cwd, "git", 10) entries, err := store.QueryLocal(context.Background(), cwd, "git", 10)
if err != nil { if err != nil {
@@ -60,9 +60,9 @@ func TestFrecencyStore_QueryGlobalDedupe(t *testing.T) {
} }
defer store.Close() defer store.Close()
_ = store.Record(context.Background(), "make build", "/repo/a") _ = store.Record(context.Background(), "make build", "/repo/a", 0)
_ = store.Record(context.Background(), "make build", "/repo/a") _ = store.Record(context.Background(), "make build", "/repo/a", 0)
_ = store.Record(context.Background(), "make build", "/repo/b") _ = store.Record(context.Background(), "make build", "/repo/b", 0)
entries, err := store.QueryGlobal(context.Background(), "make", 10) entries, err := store.QueryGlobal(context.Background(), "make", 10)
if err != nil { if err != nil {
@@ -139,7 +139,7 @@ func TestFrecencyStore_SQLiteConfigurationAndContext(t *testing.T) {
ctxCanceled, cancel := context.WithCancel(context.Background()) ctxCanceled, cancel := context.WithCancel(context.Background())
cancel() cancel()
err = store.Record(ctxCanceled, "git status", tmpDir) err = store.Record(ctxCanceled, "git status", tmpDir, 0)
if !errors.Is(err, context.Canceled) { if !errors.Is(err, context.Canceled) {
t.Errorf("expected context.Canceled from Record with canceled context, got %v", err) 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) { func TestFrecencyStore_NilReceiver(t *testing.T) {
var nilStore *FrecencyStore 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) 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 { 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) 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 BasePriority int
ContextBonus int ContextBonus int
Frecency int Frecency int
Transition int
MatchQuality int MatchQuality int
} }
@@ -25,13 +26,15 @@ type ScoreConfig struct {
WeightBasePriority float64 WeightBasePriority float64
WeightContextBonus float64 WeightContextBonus float64
WeightFrecency float64 WeightFrecency float64
WeightTransition float64
WeightMatchQuality float64 WeightMatchQuality float64
} }
var DefaultScoreConfig = ScoreConfig{ var DefaultScoreConfig = ScoreConfig{
WeightBasePriority: 0.30, WeightBasePriority: 0.30,
WeightContextBonus: 0.25, WeightContextBonus: 0.25,
WeightFrecency: 0.25, WeightFrecency: 0.15,
WeightTransition: 0.10,
WeightMatchQuality: 0.20, WeightMatchQuality: 0.20,
} }
@@ -71,11 +74,13 @@ func ScoreWithConfig(suggestions []spec.Suggestion, signals SignalSet, config Sc
bp := basePriorityFor(s) bp := basePriorityFor(s)
cb := ApplyContextRules(signals.Workspace, s.Cmd) cb := ApplyContextRules(signals.Workspace, s.Cmd)
frec := normFrec[i] frec := normFrec[i]
trans := transitionScoreFor(ExtractSkeleton(s.Cmd), signals.TransitionEntries, signals.TransitionIsLocal)
mq := matchQualityScore(s.Cmd, signals.Query) mq := matchQualityScore(s.Cmd, signals.Query)
total := config.WeightBasePriority*float64(bp) + total := config.WeightBasePriority*float64(bp) +
config.WeightContextBonus*float64(cb) + config.WeightContextBonus*float64(cb) +
config.WeightFrecency*float64(frec) + config.WeightFrecency*float64(frec) +
config.WeightTransition*float64(trans) +
config.WeightMatchQuality*float64(mq) config.WeightMatchQuality*float64(mq)
scored[i] = ScoredSuggestion{ scored[i] = ScoredSuggestion{
@@ -85,6 +90,7 @@ func ScoreWithConfig(suggestions []spec.Suggestion, signals SignalSet, config Sc
BasePriority: bp, BasePriority: bp,
ContextBonus: cb, ContextBonus: cb,
Frecency: frec, Frecency: frec,
Transition: trans,
MatchQuality: mq, MatchQuality: mq,
}, },
} }
@@ -94,6 +100,9 @@ func ScoreWithConfig(suggestions []spec.Suggestion, signals SignalSet, config Sc
if scored[i].Score != scored[j].Score { if scored[i].Score != scored[j].Score {
return 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 { if scored[i].Breakdown.Frecency != scored[j].Breakdown.Frecency {
return 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 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 { func basePriorityFor(s spec.Suggestion) int {
if s.Priority > 0 { if s.Priority > 0 {
if s.Priority > 100 { 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 'αβγδε'") 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 { type SignalSet struct {
Workspace workspace.WorkspaceInfo Workspace workspace.WorkspaceInfo
LocalFrecency []FrecencyEntry LocalFrecency []FrecencyEntry
GlobalFrecency []FrecencyEntry GlobalFrecency []FrecencyEntry
Query string TransitionEntries []TransitionEntry
RootCommand string TransitionIsLocal bool
Cwd string Query string
RootCommand string
Cwd string
} }
// CollectSignals gathers environment, workspace, and historical frecency signals for the given query and directory // 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) SignalSet { func CollectSignals(ctx context.Context, cwd, query, rootCmd string, frecency *FrecencyStore, prevCmdSkeleton string) SignalSet {
ws := workspace.DetectCached(cwd) ws := workspace.DetectCached(cwd)
if ctx == nil { if ctx == nil {
@@ -25,17 +27,25 @@ func CollectSignals(ctx context.Context, cwd, query, rootCmd string, frecency *F
} }
var local, global []FrecencyEntry var local, global []FrecencyEntry
var trans []TransitionEntry
var transIsLocal bool
if frecency != nil { if frecency != nil {
local, _ = frecency.QueryLocal(ctx, cwd, query, 50) local, _ = frecency.QueryLocal(ctx, cwd, query, 50)
global, _ = frecency.QueryGlobal(ctx, query, 50) global, _ = frecency.QueryGlobal(ctx, query, 50)
if prevCmdSkeleton != "" {
trans, transIsLocal = frecency.QueryTransitionsWithFallback(ctx, prevCmdSkeleton, cwd)
}
} }
return SignalSet{ return SignalSet{
Workspace: ws, Workspace: ws,
LocalFrecency: local, LocalFrecency: local,
GlobalFrecency: global, GlobalFrecency: global,
Query: strings.TrimSpace(query), TransitionEntries: trans,
RootCommand: strings.TrimSpace(rootCmd), TransitionIsLocal: transIsLocal,
Cwd: cwd, 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() defer store.Close()
_ = store.Record(context.Background(), "npm run dev", tmpDir) _ = store.Record(context.Background(), "npm run dev", tmpDir, 0)
_ = store.Record(context.Background(), "npm test", "/other/dir") _ = 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 { if !signals.Workspace.HasNodeProject {
t.Error("expected HasNodeProject to be true in collected signals") t.Error("expected HasNodeProject to be true in collected signals")
@@ -32,4 +33,7 @@ func TestCollectSignals(t *testing.T) {
if len(signals.GlobalFrecency) != 2 { if len(signals.GlobalFrecency) != 2 {
t.Errorf("expected global frecency to contain 2 entries, got %d", len(signals.GlobalFrecency)) 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)
}
})
}
}
+56 -1
View File
@@ -3,11 +3,13 @@ package workspace
import ( import (
"os" "os"
"path/filepath" "path/filepath"
"strings"
"sync" "sync"
) )
type WorkspaceInfo struct { type WorkspaceInfo struct {
HasGit bool HasGit bool
GitBranch string
HasNodeProject bool HasNodeProject bool
HasGoProject bool HasGoProject bool
HasRustProject bool HasRustProject bool
@@ -57,9 +59,55 @@ func Detect(cwd string) WorkspaceInfo {
} }
} }
info.HasGit, info.GitBranch = detectGitInfo(cwd)
return info return info
} }
func resolveGitHeadPath(cwd string) (hasGit bool, headPath string) {
dir := cwd
for dir != "" {
gitPath := filepath.Join(dir, ".git")
info, err := os.Stat(gitPath)
if err == nil {
if info.IsDir() {
return true, filepath.Join(gitPath, "HEAD")
}
content, errRead := os.ReadFile(gitPath)
if errRead == nil {
s := strings.TrimSpace(string(content))
if after, ok := strings.CutPrefix(s, "gitdir: "); ok {
gitDir := strings.TrimSpace(after)
if !filepath.IsAbs(gitDir) {
gitDir = filepath.Join(dir, gitDir)
}
return true, filepath.Join(gitDir, "HEAD")
}
}
return true, "" // found .git but couldn't resolve HEAD
}
parent := filepath.Dir(dir)
if parent == dir {
break
}
dir = parent
}
return false, ""
}
func detectGitInfo(cwd string) (hasGit bool, branch string) {
hasGit, headPath := resolveGitHeadPath(cwd)
if headPath != "" {
if data, errHead := os.ReadFile(headPath); errHead == nil {
s := strings.TrimSpace(string(data))
if after, ok := strings.CutPrefix(s, "ref: refs/heads/"); ok {
return hasGit, after
}
}
}
return hasGit, ""
}
type cacheEntry struct { type cacheEntry struct {
key string // cwd + "|" + dirModTime key string // cwd + "|" + dirModTime
info WorkspaceInfo info WorkspaceInfo
@@ -71,7 +119,7 @@ var (
) )
// DetectCached returns cached workspace info, invalidating when directory modtime changes // DetectCached returns cached workspace info, invalidating when directory modtime changes
// this handles mid-session file creation (e.g. go mod init) without requiring cd // or when the Git HEAD file changes (to catch branch switches that don't affect cwd modtime).
func DetectCached(cwd string) WorkspaceInfo { func DetectCached(cwd string) WorkspaceInfo {
dirInfo, err := os.Stat(cwd) dirInfo, err := os.Stat(cwd)
if err != nil { if err != nil {
@@ -79,6 +127,13 @@ func DetectCached(cwd string) WorkspaceInfo {
} }
key := cwd + "|" + dirInfo.ModTime().String() key := cwd + "|" + dirInfo.ModTime().String()
// Incorporate Git HEAD modtime into cache key to catch branch switches
if _, headPath := resolveGitHeadPath(cwd); headPath != "" {
if headInfo, err := os.Stat(headPath); err == nil {
key += "|HEAD:" + headInfo.ModTime().String()
}
}
wsCacheMu.Lock() wsCacheMu.Lock()
defer wsCacheMu.Unlock() defer wsCacheMu.Unlock()
+36
View File
@@ -4,6 +4,7 @@ import (
"os" "os"
"path/filepath" "path/filepath"
"testing" "testing"
"time"
) )
func TestDetect_GitAndGoProject(t *testing.T) { func TestDetect_GitAndGoProject(t *testing.T) {
@@ -115,3 +116,38 @@ func TestDetect_MultiEcosystems(t *testing.T) {
t.Error("expected HasK8s to be true") t.Error("expected HasK8s to be true")
} }
} }
func TestDetectCached_BranchSwitchWithoutDirChange(t *testing.T) {
tmp := t.TempDir()
gitDir := filepath.Join(tmp, ".git")
_ = os.Mkdir(gitDir, 0755)
headPath := filepath.Join(gitDir, "HEAD")
_ = os.WriteFile(headPath, []byte("ref: refs/heads/main\n"), 0644)
dirInfoBefore, _ := os.Stat(tmp)
info1 := DetectCached(tmp)
if info1.GitBranch != "main" {
t.Fatalf("expected branch 'main', got %q", info1.GitBranch)
}
// Simulate branch switch by updating .git/HEAD only
// (this does not update the modtime of tmp on typical filesystems since tmp's direct children didn't change)
_ = os.WriteFile(headPath, []byte("ref: refs/heads/feature\n"), 0644)
// Force the modtime of HEAD to be distinct to avoid flakiness on low-res file systems
infoAfter, _ := os.Stat(headPath)
newMod := infoAfter.ModTime().Add(2 * time.Second)
_ = os.Chtimes(headPath, newMod, newMod)
dirInfoAfter, _ := os.Stat(tmp)
if dirInfoBefore.ModTime() != dirInfoAfter.ModTime() {
t.Log("Note: directory modtime changed automatically on this filesystem")
}
info2 := DetectCached(tmp)
if info2.GitBranch != "feature" {
t.Fatalf("expected branch 'feature', got %q", info2.GitBranch)
}
}
+8 -1
View File
@@ -76,6 +76,13 @@ func MergeResults(query string, mode string) []spec.Suggestion {
} }
} }
if mode == "history" && normalizedQuery == "" {
if len(deduped) > maxSugg {
deduped = deduped[:maxSugg]
}
return deduped
}
if aiSugg := GetCurrentAISuggestion(); aiSugg != nil { if aiSugg := GetCurrentAISuggestion(); aiSugg != nil {
normalizedCmd := strings.TrimSpace(aiSugg.Cmd) normalizedCmd := strings.TrimSpace(aiSugg.Cmd)
if normalizedCmd != "" && normalizedCmd != normalizedQuery && strings.HasPrefix(strings.ToLower(normalizedCmd), strings.ToLower(normalizedQuery)) { if normalizedCmd != "" && normalizedCmd != normalizedQuery && strings.HasPrefix(strings.ToLower(normalizedCmd), strings.ToLower(normalizedQuery)) {
@@ -105,7 +112,7 @@ func MergeResults(query string, mode string) []spec.Suggestion {
ctxTimeout, cancel := context.WithTimeout(context.Background(), 500*time.Millisecond) ctxTimeout, cancel := context.WithTimeout(context.Background(), 500*time.Millisecond)
defer cancel() defer cancel()
store, _ := scoring.GetFrecencyStore() store, _ := scoring.GetFrecencyStore()
signals := scoring.CollectSignals(ctxTimeout, cwd, query, rootCmd, store) signals := scoring.CollectSignals(ctxTimeout, cwd, query, rootCmd, store, getPrevSkeleton())
scored := scoring.Score(deduped, signals) scored := scoring.Score(deduped, signals)
finalResults := make([]spec.Suggestion, 0, len(scored)) finalResults := make([]spec.Suggestion, 0, len(scored))
+29 -1
View File
@@ -20,7 +20,6 @@ func TestMergeResults(t *testing.T) {
} }
}) })
t.Run("Limit 100", func(t *testing.T) { t.Run("Limit 100", func(t *testing.T) {
res := MergeResults("a", "history") res := MergeResults("a", "history")
if len(res) > 100 { if len(res) > 100 {
@@ -50,3 +49,32 @@ func TestMergeResults(t *testing.T) {
} }
}) })
} }
func TestPrevRecordedCommandState(t *testing.T) {
origRegistry := spec.Registry
t.Cleanup(func() {
spec.Registry = origRegistry
})
spec.ResetRegistry()
spec.Register(&spec.Spec{
Name: "git",
Subcommands: []spec.Subcommand{
{Name: "checkout"},
},
})
setPrevRecordedInfo("git checkout feature-abc", "/repo/dir")
skel, cwd := getPrevRecordedInfo()
if skel != "git checkout" || cwd != "/repo/dir" {
t.Errorf("expected skel='git checkout' and cwd='/repo/dir', got skel=%q, cwd=%q", skel, cwd)
}
if gotSkel := getPrevSkeleton(); gotSkel != "git checkout" {
t.Errorf("expected getPrevSkeleton()='git checkout', got %q", gotSkel)
}
setPrevRecordedInfo("", "")
if gotSkel := getPrevSkeleton(); gotSkel != "" {
t.Errorf("expected empty getPrevSkeleton() after reset, got %q", gotSkel)
}
}
+63 -5
View File
@@ -10,6 +10,7 @@ import (
"os/exec" "os/exec"
"os/signal" "os/signal"
"path/filepath" "path/filepath"
"strconv"
"strings" "strings"
"sync" "sync"
"sync/atomic" "sync/atomic"
@@ -28,6 +29,37 @@ import (
"golang.org/x/term" "golang.org/x/term"
) )
var (
prevRecordedCommand string
prevCmdCwd string
prevCmdMu sync.Mutex
)
func getPrevSkeleton() string {
prevCmdMu.Lock()
defer prevCmdMu.Unlock()
if prevRecordedCommand == "" {
return ""
}
return scoring.ExtractSkeleton(prevRecordedCommand)
}
func getPrevRecordedInfo() (string, string) {
prevCmdMu.Lock()
defer prevCmdMu.Unlock()
if prevRecordedCommand == "" {
return "", ""
}
return scoring.ExtractSkeleton(prevRecordedCommand), prevCmdCwd
}
func setPrevRecordedInfo(cmd, cwd string) {
prevCmdMu.Lock()
defer prevCmdMu.Unlock()
prevRecordedCommand = cmd
prevCmdCwd = cwd
}
func loadMode() string { func loadMode() string {
mode := config.Get().Core.Mode mode := config.Get().Core.Mode
if mode == "last" { if mode == "last" {
@@ -322,7 +354,13 @@ func runWrapper() {
continue continue
} }
if query == "IRIS_CMD_STOP" { if query == "IRIS_CMD_STOP" || strings.HasPrefix(query, "IRIS_CMD_STOP:") {
exitCode := 0
if after, ok := strings.CutPrefix(query, "IRIS_CMD_STOP:"); ok {
if code, err := strconv.Atoi(after); err == nil {
exitCode = code
}
}
isCommandActive.Store(false) isCommandActive.Store(false)
SetCurrentAISuggestion(nil) SetCurrentAISuggestion(nil)
bufferMu.Lock() bufferMu.Lock()
@@ -330,8 +368,11 @@ func runWrapper() {
lastSubmittedCommand = "" lastSubmittedCommand = ""
bufferMu.Unlock() bufferMu.Unlock()
if cmdToRecord != "" { if cmdToRecord != "" {
integration.RecordSessionCommand(cmdToRecord)
cwd := spec.GetCWD() cwd := spec.GetCWD()
go func(c, d string) { prevSkeleton, prevCwd := getPrevRecordedInfo()
currSkeleton := scoring.ExtractSkeleton(cmdToRecord)
go func(c, d string, code int, pSkel, pCwd, cSkel string) {
defer func() { defer func() {
if r := recover(); r != nil { if r := recover(); r != nil {
WriteCrashLog(r) WriteCrashLog(r)
@@ -340,9 +381,13 @@ func runWrapper() {
ctxRecord, cancel := context.WithTimeout(context.Background(), 1500*time.Millisecond) ctxRecord, cancel := context.WithTimeout(context.Background(), 1500*time.Millisecond)
defer cancel() defer cancel()
if store, err := scoring.GetFrecencyStore(); err == nil && store != nil { if store, err := scoring.GetFrecencyStore(); err == nil && store != nil {
_ = store.Record(ctxRecord, c, d) _ = store.Record(ctxRecord, c, d, code)
if pSkel != "" && cSkel != "" {
_ = store.RecordTransition(ctxRecord, pSkel, cSkel, d, code)
}
} }
}(cmdToRecord, cwd) }(cmdToRecord, cwd, exitCode, prevSkeleton, prevCwd, currSkeleton)
setPrevRecordedInfo(cmdToRecord, cwd)
} }
// hook: after user executes a command, print the update notice exactly once per session // hook: after user executes a command, print the update notice exactly once per session
if !updatePrinted { if !updatePrinted {
@@ -605,17 +650,30 @@ func runWrapper() {
if inputSlice[i+2] == 'A' { if inputSlice[i+2] == 'A' {
arrowDir = "up" arrowDir = "up"
} }
moved, _ := overlay.MoveCursor(arrowDir) moved, selectedCmd := overlay.MoveCursor(arrowDir)
if !moved { if !moved {
i += 2 i += 2
continue continue
} }
bufferMu.Lock() bufferMu.Lock()
activeModeMu.RLock()
isHistMode := activeMode == "history"
activeModeMu.RUnlock()
var toWrite []byte
if isHistMode && selectedCmd != "" {
naiveBuffer = selectedCmd
cursorOffset = 0
toWrite = append([]byte{0x15}, selectedCmd...)
}
bufCopy := naiveBuffer bufCopy := naiveBuffer
offsetCopy := cursorOffset offsetCopy := cursorOffset
bufferMu.Unlock() bufferMu.Unlock()
if len(toWrite) > 0 {
_, _ = ptmx.Write(toWrite)
}
var b strings.Builder var b strings.Builder
if !disableGhostText.Load() { if !disableGhostText.Load() {
b.WriteString(overlay.RenderGhostText(bufCopy, true, offsetCopy == 0)) b.WriteString(overlay.RenderGhostText(bufCopy, true, offsetCopy == 0))
+87
View File
@@ -0,0 +1,87 @@
package spec
import (
"slices"
"strings"
)
// TryExtractSkeleton walks the registered subcommand tree for the given command buffer.
// It stops when a token doesn't match any registered subcommand name (treating that token as an argument/value)
// or when encountering a flag ('-' prefix).
// Returns the normalized subcommand skeleton (e.g. "git checkout feature-x" -> "git checkout") and true if a spec exists.
func TryExtractSkeleton(buf string) (string, bool) {
buf = strings.TrimSpace(buf)
if buf == "" {
return "", false
}
aliases := GetAliasesCopy()
tokens := Tokenize(buf)
// filter out empty tokens (e.g. from trailing space)
var cleanTokens []string
for _, t := range tokens {
if t != "" {
cleanTokens = append(cleanTokens, t)
}
}
if len(cleanTokens) == 0 {
return "", false
}
// expand shell alias for the root command if present
if target, ok := aliases[cleanTokens[0]]; ok {
aliasTokens := Tokenize(target)
var cleanAliasTokens []string
for _, t := range aliasTokens {
if t != "" {
cleanAliasTokens = append(cleanAliasTokens, t)
}
}
cleanTokens = append(cleanAliasTokens, cleanTokens[1:]...)
}
if len(cleanTokens) == 0 {
return "", false
}
rootName := cleanTokens[0]
spec, exists := Registry[rootName]
if !exists || spec == nil {
return "", false
}
skeletonTokens := []string{rootName}
currentSubs := spec.Subcommands
for i := 1; i < len(cleanTokens); i++ {
tok := cleanTokens[i]
if strings.HasPrefix(tok, "-") || strings.Contains(tok, "=") {
continue
}
found := false
for _, sub := range currentSubs {
if sub.Name == tok || slices.Contains(sub.Aliases, tok) {
skeletonTokens = append(skeletonTokens, sub.Name)
currentSubs = sub.Subcommands
found = true
break
}
}
if found {
continue
}
// If the previous token was a flag (started with '-'), this token might be the flag's argument/value (e.g. -C /tmp).
// In that case, we skip this token and continue looking for subcommands in subsequent tokens.
if i > 0 && strings.HasPrefix(cleanTokens[i-1], "-") {
continue
}
// token did not match any registered subcommand and is not a flag argument -> must be a positional argument (branch, file, etc.)
break
}
return strings.Join(skeletonTokens, " "), true
}
+97
View File
@@ -0,0 +1,97 @@
package spec
import (
"testing"
)
func TestTryExtractSkeleton(t *testing.T) {
ResetRegistry()
Register(&Spec{
Name: "git",
Subcommands: []Subcommand{
{
Name: "checkout",
Aliases: []string{"co"},
},
{
Name: "remote",
Subcommands: []Subcommand{
{Name: "add"},
},
},
},
})
Register(&Spec{
Name: "docker",
Subcommands: []Subcommand{
{
Name: "compose",
Subcommands: []Subcommand{
{Name: "up"},
},
},
},
})
tests := []struct {
name string
input string
wantSkeleton string
wantOk bool
}{
{
name: "git checkout with argument",
input: "git checkout feature-x -b",
wantSkeleton: "git checkout",
wantOk: true,
},
{
name: "git co alias with argument",
input: "git co feature-y",
wantSkeleton: "git checkout",
wantOk: true,
},
{
name: "git remote add with arguments",
input: "git remote add origin https://github.com/test/test.git",
wantSkeleton: "git remote add",
wantOk: true,
},
{
name: "git global flag before subcommand",
input: "git -C /tmp checkout main",
wantSkeleton: "git checkout",
wantOk: true,
},
{
name: "docker compose up with flags",
input: "docker compose up -d",
wantSkeleton: "docker compose up",
wantOk: true,
},
{
name: "unregistered spec",
input: "cargo build --release",
wantSkeleton: "",
wantOk: false,
},
{
name: "empty input",
input: " ",
wantSkeleton: "",
wantOk: false,
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
got, ok := TryExtractSkeleton(tc.input)
if ok != tc.wantOk {
t.Fatalf("TryExtractSkeleton(%q) ok = %v, want %v", tc.input, ok, tc.wantOk)
}
if got != tc.wantSkeleton {
t.Errorf("TryExtractSkeleton(%q) = %q, want %q", tc.input, got, tc.wantSkeleton)
}
})
}
}